package main

import (
	"context"
	"encoding/json"
	"fmt"
	"io"
	"net"
	"net/http"
	"os"
	"os/exec"
	"path/filepath"
	"strings"
	"syscall"
	"time"
)

// Black-box operational checks, including a real SIGTERM with a database-blocked
// request in flight. Only a marked task database can be used.
func checkOperations(ctx context.Context, dsn, binary, output string) error {
	c, err := openDB(ctx, dsn)
	if err != nil {
		return err
	}
	defer c.Close(ctx)
	if err = fixture(ctx, c); err != nil {
		return err
	}
	binary, err = filepath.Abs(binary)
	if err != nil {
		return err
	}
	if err = os.MkdirAll(filepath.Dir(output), 0755); err != nil {
		return err
	}
	log, err := os.Create(output + ".log")
	if err != nil {
		return err
	}
	defer log.Close()
	listener, err := net.Listen("tcp", "127.0.0.1:0")
	if err != nil {
		return err
	}
	address := listener.Addr().String()
	listener.Close()
	cmd := exec.Command(binary)
	cmd.Env = append(os.Environ(), "DATABASE_URL="+dsn, "APP_PROFILE=production", "ALLOWED_HOSTS=127.0.0.1", "LISTEN_ADDR="+address, "SQL_TRACE=")
	cmd.Stdout = log
	cmd.Stderr = log
	if err = cmd.Start(); err != nil {
		return err
	}
	done := make(chan error, 1)
	go func() { done <- cmd.Wait() }()
	stopped := false
	defer func() {
		if !stopped {
			cmd.Process.Signal(syscall.SIGTERM)
			select {
			case <-done:
			case <-time.After(12 * time.Second):
				cmd.Process.Kill()
				<-done
			}
		}
	}()
	client := &http.Client{Timeout: 20 * time.Second}
	defer client.CloseIdleConnections()
	base := "http://" + address
	ready := false
	for i := 0; i < 100; i++ {
		resp, e := client.Get(base + "/healthz")
		if e == nil {
			resp.Body.Close()
			if resp.StatusCode == 200 {
				ready = true
				break
			}
		}
		select {
		case e := <-done:
			stopped = true
			return fmt.Errorf("server exited during startup: %v; see log", e)
		case <-time.After(50 * time.Millisecond):
		}
	}
	if !ready {
		return fmt.Errorf("server did not become healthy")
	}
	checks := []string{}
	check := func(name, method, path, body, host string, status int, chunked bool) error {
		var reader io.Reader = strings.NewReader(body)
		if chunked {
			reader = io.NopCloser(reader)
		}
		r, e := http.NewRequestWithContext(ctx, method, base+path, reader)
		if e != nil {
			return e
		}
		if host != "" {
			r.Host = host
			r.Header.Set("X-Forwarded-Host", "127.0.0.1")
		}
		r.Header.Set("Content-Type", "application/json")
		resp, e := client.Do(r)
		if e != nil {
			return e
		}
		defer resp.Body.Close()
		raw, e := io.ReadAll(resp.Body)
		if e != nil {
			return e
		}
		if resp.StatusCode != status {
			return fmt.Errorf("%s: status %d, want %d", name, resp.StatusCode, status)
		}
		if resp.Header.Get("X-Request-ID") == "" {
			return fmt.Errorf("%s: missing request ID", name)
		}
		if strings.Contains(string(raw), "Traceback") {
			return fmt.Errorf("%s: leaked legacy debug traceback", name)
		}
		checks = append(checks, name)
		return nil
	}
	for _, tc := range []struct {
		name, method, path, body, host string
		status                         int
		chunked                        bool
	}{
		{"liveness", "GET", "/healthz", "", "", 200, false},
		{"readiness", "GET", "/readyz", "", "", 200, false},
		{"host policy ignores proxy spoof", "GET", "/healthz", "", "attacker.invalid", 400, false},
		{"successful user read", "GET", "/api/users/1", "", "", 200, false},
		{"duplicate error suppressed", "GET", "/api/users/find?email=duplicate%40example.com", "", "", 500, false},
		{"partial failure suppressed", "POST", "/api/posts", `{"author_id":1,"title":"Partial","body":"X","tag_slugs":["python","missing"]}`, "", 500, false},
		{"known-length body limit", "POST", "/api/posts", strings.Repeat("x", (1<<20)+1), "", 413, false},
		{"chunked body limit", "POST", "/api/posts", strings.Repeat("x", (1<<20)+1), "", 413, true},
	} {
		if err = check(tc.name, tc.method, tc.path, tc.body, tc.host, tc.status, tc.chunked); err != nil {
			return err
		}
	}
	var partial bool
	if err = c.QueryRow(ctx, `SELECT EXISTS(SELECT 1 FROM blog_post p WHERE p.id=5 AND p.title='Partial') AND (SELECT count(*) FROM blog_post_tags WHERE post_id=5)=1`).Scan(&partial); err != nil {
		return err
	}
	if !partial {
		return fmt.Errorf("production changed partial-write behavior")
	}
	checks = append(checks, "partial database effects preserved")
	tx, err := c.Begin(ctx)
	if err != nil {
		return err
	}
	defer tx.Rollback(ctx)
	if _, err = tx.Exec(ctx, "UPDATE blog_post SET view_count=view_count WHERE id=1"); err != nil {
		return err
	}
	requestDone := make(chan error, 1)
	go func() {
		resp, e := client.Get(base + "/api/posts/1")
		if e == nil {
			io.Copy(io.Discard, resp.Body)
			resp.Body.Close()
			if resp.StatusCode != 200 {
				e = fmt.Errorf("in-flight status %d", resp.StatusCode)
			}
		}
		requestDone <- e
	}()
	observer, err := openDB(ctx, dsn)
	if err != nil {
		return err
	}
	defer observer.Close(ctx)
	blocked := false
	for i := 0; i < 100; i++ {
		if err = observer.QueryRow(ctx, `SELECT EXISTS(SELECT 1 FROM pg_stat_activity WHERE datname=current_database() AND wait_event_type='Lock' AND query LIKE 'UPDATE blog_post SET title=%')`).Scan(&blocked); err != nil {
			return err
		}
		if blocked {
			break
		}
		time.Sleep(20 * time.Millisecond)
	}
	if !blocked {
		return fmt.Errorf("could not establish in-flight database wait")
	}
	if err = cmd.Process.Signal(syscall.SIGTERM); err != nil {
		return err
	}
	// Confirm SIGTERM does not end the process before its blocked handler finishes.
	select {
	case e := <-done:
		stopped = true
		return fmt.Errorf("exited before draining: %v", e)
	case <-time.After(100 * time.Millisecond):
	}
	if err = tx.Commit(ctx); err != nil {
		return err
	}
	if err = <-requestDone; err != nil {
		return err
	}
	select {
	case err = <-done:
		stopped = true
		if err != nil {
			return err
		}
	case <-time.After(12 * time.Second):
		return fmt.Errorf("shutdown deadline exceeded")
	}
	checks = append(checks, "SIGTERM drained blocked request and exited zero")
	raw, err := json.MarshalIndent(map[string]any{"passed": true, "checks": checks, "captured_at": time.Now().UTC()}, "", "  ")
	if err != nil {
		return err
	}
	return os.WriteFile(output, append(raw, '\n'), 0644)
}
