"""Tracer tests: global db identity and legacy migration."""

from __future__ import annotations

import json
import sqlite3
import sys
import threading
from pathlib import Path

import pytest

from adw_modules import git_helper
from adw_modules.data_types import ConfigDefaults, EventRecord, ObservabilityConfig, SSSFConfig
from adw_modules.session import ensure
from adw_modules.tracer import Tracer, connect_db

LEGACY_SCHEMA = """
CREATE TABLE sessions (
  adw_id        TEXT PRIMARY KEY,
  request       TEXT,
  status        TEXT,
  engineer      TEXT,
  started_at    TEXT, ended_at TEXT,
  total_tokens  INTEGER, total_cost REAL
);
CREATE TABLE agent_attempts (
  attempt_id TEXT PRIMARY KEY,
  adw_id TEXT,
  phase_id TEXT,
  parent_id TEXT,
  agent TEXT,
  session_id TEXT,
  command TEXT,
  started_at TEXT NOT NULL
);
CREATE TABLE processes (
  id INTEGER PRIMARY KEY AUTOINCREMENT,
  adw_id TEXT,
  kind TEXT,
  name TEXT,
  pid INTEGER,
  command TEXT,
  started_at TEXT,
  ended_at TEXT
);
"""


def _pragma_values(conn: sqlite3.Connection) -> dict[str, int | str]:
    return {
        "journal_mode": conn.execute("PRAGMA journal_mode").fetchone()[0],
        "synchronous": conn.execute("PRAGMA synchronous").fetchone()[0],
        "busy_timeout": conn.execute("PRAGMA busy_timeout").fetchone()[0],
    }


def test_connect_db_applies_wal_pragmas(tmp_path: Path) -> None:
    conn = connect_db(tmp_path / "sssf.db")
    try:
        assert _pragma_values(conn) == {
            "journal_mode": "wal",
            "synchronous": 1,          # PRAGMA synchronous=NORMAL
            "busy_timeout": 5000,
        }
    finally:
        conn.close()


def test_factory_run_keeps_one_explicit_request_identity(tmp_path: Path) -> None:
    tracer = Tracer(tmp_path / "sssf.db", tmp_path / "events.jsonl")
    try:
        tracer.session_start("adw-linked", "alice", request_id="manual-request-42")
        assert tracer.conn.execute(
            "SELECT request_id, adw_id FROM request_run_links"
        ).fetchall() == [("manual-request-42", "adw-linked")]

        tracer.session_start("adw-linked", "alice", request_id="manual-request-42")
        with pytest.raises(ValueError, match="different request"):
            tracer.session_start("adw-linked", "alice", request_id="manual-other")
        assert tracer.conn.execute(
            "SELECT request_id FROM request_run_links WHERE adw_id='adw-linked'"
        ).fetchone() == ("manual-request-42",)
    finally:
        tracer.conn.close()


def test_two_repos_share_one_db(tmp_path: Path) -> None:
    db = tmp_path / "sssf.db"
    events = tmp_path / "events.jsonl"
    repo_a = tmp_path / "project-a"
    repo_b = tmp_path / "project-b"
    repo_a.mkdir()
    repo_b.mkdir()

    tracer_a = Tracer(db, events)
    try:
        tracer_a.session_start("adw_alpha", "alice", repo=repo_a)
        tracer_a.session_finish("adw_alpha", ok=True)
    finally:
        tracer_a.conn.close()

    tracer_b = Tracer(db, events)
    try:
        tracer_b.session_start("adw_beta", "bob", repo=repo_b)
        tracer_b.session_finish("adw_beta", ok=True)
    finally:
        tracer_b.conn.close()

    reader = connect_db(db)
    try:
        rows = reader.execute(
            "SELECT adw_id, repo FROM sessions ORDER BY adw_id"
        ).fetchall()
        assert rows == [
            ("adw_alpha", str(repo_a.resolve())),
            ("adw_beta", str(repo_b.resolve())),
        ]
        assert _pragma_values(reader)["journal_mode"] == "wal"
    finally:
        reader.close()


def test_legacy_db_without_repo_column_is_readable(tmp_path: Path) -> None:
    db = tmp_path / "legacy.db"
    events = tmp_path / "events.jsonl"

    conn = sqlite3.connect(db, isolation_level=None)
    try:
        conn.executescript(LEGACY_SCHEMA)
        conn.execute(
            "INSERT INTO sessions VALUES (?,?,?,?,?,?,?,?)",
            ("legacy1", "fix bug", "success", "alice",
             "2026-01-01T00:00:00Z", "2026-01-01T01:00:00Z", 100, 0.01),
        )
    finally:
        conn.close()

    tracer = Tracer(db, events)
    try:
        row = tracer.conn.execute(
            "SELECT adw_id, request, status, engineer, repo FROM sessions WHERE adw_id=?",
            ("legacy1",),
        ).fetchone()
        columns = {name for _, name, *_ in tracer.conn.execute("PRAGMA table_info(sessions)")}
        attempt_columns = {
            name for _, name, *_ in tracer.conn.execute("PRAGMA table_info(agent_attempts)")
        }
        process_columns = {
            name for _, name, *_ in tracer.conn.execute("PRAGMA table_info(processes)")
        }
    finally:
        tracer.conn.close()

    assert row == ("legacy1", "fix bug", "success", "alice", None)
    assert "repo" in columns
    assert "host" in columns
    assert {"host", "account", "model", "timeout_kind"} <= attempt_columns
    assert "start_ticks" in process_columns


def test_fresh_db_has_execution_identity_columns(tmp_path: Path) -> None:
    tracer = Tracer(tmp_path / "fresh.db", tmp_path / "events.jsonl")
    try:
        session_columns = {
            name for _, name, *_ in tracer.conn.execute("PRAGMA table_info(sessions)")
        }
        attempt_columns = {
            name for _, name, *_ in tracer.conn.execute("PRAGMA table_info(agent_attempts)")
        }
        process_columns = {
            name for _, name, *_ in tracer.conn.execute("PRAGMA table_info(processes)")
        }
    finally:
        tracer.conn.close()

    assert "host" in session_columns
    assert {"host", "account", "model", "timeout_kind"} <= attempt_columns
    assert "start_ticks" in process_columns


def test_session_ensure_records_canonical_repo(tmp_path: Path, monkeypatch) -> None:
    """session.ensure is the real lifecycle seam — not a direct Tracer call."""
    repo = tmp_path / "target-repo"
    repo.mkdir()
    monkeypatch.setattr(git_helper, "repo_root", lambda: repo)

    db = tmp_path / "sssf.db"
    cfg = SSSFConfig(
        defaults=ConfigDefaults(data_dir=str(tmp_path / "data")),
        observability=ObservabilityConfig(db=str(db)),
    )
    monkeypatch.setattr(sys, "argv", ["adw_test"])
    monkeypatch.setenv("FACTORY_REQUEST_ID", "manual-integration-request")

    run = ensure(cfg, adw_id="integration123")
    try:
        row = run.tracer.conn.execute(
            "SELECT repo FROM sessions WHERE adw_id=?",
            ("integration123",),
        ).fetchone()
        assert row == (str(repo.resolve()),)
        assert run.tracer.conn.execute(
            "SELECT request_id FROM request_run_links WHERE adw_id=?",
            ("integration123",),
        ).fetchone() == ("manual-integration-request",)
    finally:
        run.tracer.conn.close()


def test_session_start_on_conflict_preserves_and_fills_repo(tmp_path: Path) -> None:
    db = tmp_path / "sssf.db"
    events = tmp_path / "events.jsonl"
    repo_a = tmp_path / "project-a"
    repo_b = tmp_path / "project-b"
    repo_a.mkdir()
    repo_b.mkdir()

    tracer = Tracer(db, events)
    try:
        tracer.session_start("joined", "alice", repo=repo_a)
        tracer.session_start("joined", "alice", adw_name="adw_followup")
        row = tracer.conn.execute(
            "SELECT repo FROM sessions WHERE adw_id=?", ("joined",),
        ).fetchone()
        assert row == (str(repo_a.resolve()),)

        tracer.conn.execute("DELETE FROM sessions WHERE adw_id=?", ("joined",))
        tracer.session_start("joined", "alice")
        tracer.session_start("joined", "alice", repo=repo_b, adw_name="adw_followup")
        row = tracer.conn.execute(
            "SELECT repo FROM sessions WHERE adw_id=?", ("joined",),
        ).fetchone()
        assert row == (str(repo_b.resolve()),)
    finally:
        tracer.conn.close()


def test_concurrent_tracer_constructors_migrate_legacy_db(tmp_path: Path) -> None:
    db = tmp_path / "legacy.db"
    events = tmp_path / "events.jsonl"

    conn = sqlite3.connect(db, isolation_level=None)
    try:
        conn.executescript(LEGACY_SCHEMA)
    finally:
        conn.close()

    barrier = threading.Barrier(2)
    errors: list[BaseException] = []

    def open_tracer() -> None:
        try:
            barrier.wait(timeout=5)
            Tracer(db, events).conn.close()
        except BaseException as exc:
            errors.append(exc)

    threads = [threading.Thread(target=open_tracer) for _ in range(2)]
    for thread in threads:
        thread.start()
    for thread in threads:
        thread.join(timeout=10)

    assert not errors
    reader = connect_db(db)
    try:
        columns = {name for _, name, *_ in reader.execute("PRAGMA table_info(sessions)")}
        assert "repo" in columns
    finally:
        reader.close()


def test_timeout_kind_migration_is_idempotent_on_existing_db(tmp_path: Path) -> None:
    db = tmp_path / "legacy.db"
    events = tmp_path / "events.jsonl"
    conn = sqlite3.connect(db, isolation_level=None)
    try:
        conn.executescript(LEGACY_SCHEMA)
    finally:
        conn.close()

    Tracer(db, events).conn.close()
    Tracer(db, events).conn.close()

    reader = sqlite3.connect(db)
    try:
        columns = [row[1] for row in reader.execute("PRAGMA table_info(agent_attempts)")]
        assert columns.count("timeout_kind") == 1
    finally:
        reader.close()


def test_phase_diff_linkage_migration_is_idempotent_on_existing_db(tmp_path: Path) -> None:
    db = tmp_path / "legacy.db"
    events = tmp_path / "events.jsonl"
    conn = sqlite3.connect(db, isolation_level=None)
    try:
        conn.executescript("""
            CREATE TABLE phase_diffs (
              id INTEGER PRIMARY KEY AUTOINCREMENT,
              adw_id TEXT NOT NULL,
              phase_id TEXT NOT NULL,
              attempt INTEGER,
              files_json TEXT,
              insertions INTEGER,
              deletions INTEGER,
              diff_text TEXT,
              truncated INTEGER,
              created_at TEXT NOT NULL
            );
        """)
    finally:
        conn.close()

    Tracer(db, events).conn.close()
    Tracer(db, events).conn.close()

    reader = sqlite3.connect(db)
    try:
        columns = [row[1] for row in reader.execute("PRAGMA table_info(phase_diffs)")]
        phase_columns = [row[1] for row in reader.execute("PRAGMA table_info(phases)")]
        assert phase_columns.count("task_id") == 1
        assert columns.count("task_id") == 1
        assert columns.count("attempt_id") == 1
    finally:
        reader.close()


def test_tool_calls_open_and_close_in_sequence_order(tmp_path: Path) -> None:
    tracer = Tracer(tmp_path / "trace.db", tmp_path / "events.jsonl")
    try:
        tracer.tool_call_start("call-2", "attempt-1", 1, "read", {"path": "b"},
                               "2026-08-08T10:00:01Z")
        tracer.tool_call_start("call-1", "attempt-1", 0, "ls", {"path": "a"},
                               "2026-08-08T10:00:00Z")
        tracer.tool_call_finish("call-1", ended_at="2026-08-08T10:00:00.125Z",
                                duration_ms=125, ok=True, result="a.txt")
        rows = tracer.conn.execute(
            "SELECT tool_call_id,seq,tool_name,args_json,ended_at,duration_ms,ok,result_excerpt"
            " FROM tool_calls WHERE attempt_id=? ORDER BY seq", ("attempt-1",),
        ).fetchall()
    finally:
        tracer.conn.close()

    assert rows == [
        ("call-1", 0, "ls", '{"path": "a"}', "2026-08-08T10:00:00.125Z",
         125, 1, "a.txt"),
        ("call-2", 1, "read", '{"path": "b"}', None, None, None, None),
    ]


def test_tool_result_excerpt_is_explicitly_truncated(tmp_path: Path) -> None:
    tracer = Tracer(tmp_path / "trace.db", tmp_path / "events.jsonl")
    try:
        tracer.tool_call_start("call-1", "attempt-1", 0, "ls", {})
        tracer.tool_call_finish("call-1", ended_at=None, duration_ms=1, ok=True,
                                result="x" * 3000)
        excerpt = tracer.conn.execute(
            "SELECT result_excerpt FROM tool_calls WHERE tool_call_id='call-1'"
        ).fetchone()[0]
    finally:
        tracer.conn.close()

    assert len(excerpt) <= 2000
    assert excerpt.endswith("… [truncated]")


def test_kubernetes_trace_emits_only_allowlisted_lifecycle_metadata(
    tmp_path: Path, monkeypatch, capsys,
) -> None:
    attempt = "a" * 24
    monkeypatch.setenv("FACTORY_ATTEMPT_ENVELOPE", json.dumps({"attempt_id": attempt}))
    trace = Tracer(tmp_path / "trace.db", tmp_path / "events.jsonl")
    try:
        trace.session_start(
            "adw-safe", "factory", adw_name="adw_test", repo=tmp_path, preset="k3s",
        )
        parent = trace.event(EventRecord(
            adw_id="adw-safe", phase_id="phase-1", type="agent_start", name="builder",
            payload={"model": "model-1", "session_id": "session-1", "purpose": "PROMPT-SECRET"},
        ))
        agent_attempt = trace.agent_attempt_start(
            "adw-safe", "phase-1", "builder", "session-1",
            "tool --token TOP-SECRET", host="worker-pod", account="account-1",
            provider="provider", model="model-1", system_prompt="SYSTEM-SECRET",
            user_prompt="USER-SECRET", parent_id=parent,
        )
        trace.event(EventRecord(
            adw_id="adw-safe", phase_id="phase-1", type="tool_call_start",
            name="read: /workspace/ARG-NAME-SECRET",
            payload={"tool_call_id": "tool-1", "attempt_id": agent_attempt, "seq": 1,
                     "tool": "read", "args": {"path": "ARG-SECRET"}},
        ))
        trace.event(EventRecord(
            adw_id="adw-safe", phase_id="phase-1", type="tool_call",
            name="read: /workspace/RESULT-NAME-SECRET",
            payload={"tool_call_id": "tool-1", "attempt_id": agent_attempt,
                     "duration_ms": 2, "ok": True, "result_snippet": "RESULT-SECRET"},
        ))
        trace.event(EventRecord(
            adw_id="adw-safe", phase_id="phase-1", type="tool_call_start",
            name="shell: MALFORMED-TOOL-NAME-SECRET",
            payload={"tool_call_id": "unmirrored-tool", "attempt_id": agent_attempt},
        ))
        trace.event(EventRecord(
            adw_id="adw-safe", phase_id="phase-1", type="tool_call",
            name="shell: MALFORMED-RESULT-NAME-SECRET",
            payload={"tool_call_id": "unmirrored-tool", "attempt_id": agent_attempt,
                     "duration_ms": 1, "ok": True},
        ))
        trace.event(EventRecord(
            adw_id="adw-safe", phase_id="phase-1", type="tool_call_start",
            name="shell: UNFINISHED-TOOL-NAME-SECRET",
            payload={"tool_call_id": "unfinished-tool", "attempt_id": agent_attempt,
                     "seq": 2, "tool": "shell"},
        ))
        trace.event(EventRecord(
            adw_id="adw-safe", phase_id="phase-1", type="tool_call_start",
            name="shell: MISMATCHED-TOOL-NAME-SECRET",
            payload={"tool_call_id": "mismatched-tool", "attempt_id": agent_attempt,
                     "seq": 3, "tool": "shell"},
        ))
        trace.event(EventRecord(
            adw_id="adw-safe", phase_id="phase-1", type="tool_call",
            name="shell: MISMATCHED-RESULT-NAME-SECRET",
            payload={"tool_call_id": "mismatched-tool", "attempt_id": "other-attempt",
                     "duration_ms": 1, "ok": True},
        ))
        trace.event(EventRecord(
            adw_id="adw-safe", phase_id="phase-1", type="error", name="builder",
            payload={"error": "EVENT-ERROR-SECRET"},
        ))
        trace.agent_attempt_finish(
            agent_attempt, returncode=1, signal=None, timed_out=False,
            stderr_path="/secret/stderr", tokens=7,
            usage={"input_tokens": 3, "output_tokens": 4, "cost": 0.5,
                   "credential": "USAGE-SECRET"},
            error="ERROR-SECRET", provider_failure={"payload": "PROVIDER-SECRET"},
        )
        trace.session_finish("adw-safe", ok=False)
        trace.event(EventRecord(
            adw_id="adw-safe", type="log", name="console",
            payload={"message": "AFTER-FINISH-SECRET", "level": "error"},
        ))
    finally:
        trace.conn.close()

    output = capsys.readouterr().out
    assert all(secret not in output for secret in [
        "TOP-SECRET", "SYSTEM-SECRET", "USER-SECRET", "USAGE-SECRET",
        "ERROR-SECRET", "PROVIDER-SECRET", "/secret/stderr", "PROMPT-SECRET",
        "ARG-SECRET", "RESULT-SECRET", "ARG-NAME-SECRET", "RESULT-NAME-SECRET",
        "MALFORMED-TOOL-NAME-SECRET", "MALFORMED-RESULT-NAME-SECRET",
        "UNFINISHED-TOOL-NAME-SECRET", "MISMATCHED-TOOL-NAME-SECRET",
        "MISMATCHED-RESULT-NAME-SECRET", "EVENT-ERROR-SECRET",
        "AFTER-FINISH-SECRET",
    ])
    records = [
        json.loads(line.removeprefix("FACTORY_K3S_TRACE_V1 "))
        for line in output.splitlines()
    ]
    assert [record["kind"] for record in records] == [
        "session_start", "event", "agent_attempt_start", "event", "event", "event",
        "agent_attempt_finish", "session_finish",
    ]
    assert all(record["attempt_id"] == attempt and record["version"] == 1 for record in records)
    assert records[2]["model"] == "model-1" and "command" not in records[2]
    assert records[3]["metadata"] == {
        "tool_call_id": "tool-1", "attempt_id": agent_attempt, "tool": "read", "seq": 1,
    }
    assert "name" not in records[3]
    assert records[4]["metadata"] == {
        "tool_call_id": "tool-1", "attempt_id": agent_attempt,
        "duration_ms": 2, "ok": True,
    }
    assert "name" not in records[4]
    assert records[5]["metadata"] == {"error": True}
    assert records[6]["input_tokens"] == 3 and records[6]["output_tokens"] == 4


def test_local_tracing_does_not_emit_kubernetes_records(tmp_path: Path, capsys) -> None:
    trace = Tracer(tmp_path / "trace.db", tmp_path / "events.jsonl")
    try:
        trace.session_start("local-run", "factory", repo=tmp_path)
        trace.session_finish("local-run", ok=True)
    finally:
        trace.conn.close()
    assert capsys.readouterr().out == ""
