import gc
import json
import os
import re
import sys
import threading
import time
import traceback
from http.server import BaseHTTPRequestHandler, ThreadingHTTPServer
from pathlib import Path
from urllib.parse import parse_qsl, urlparse

sys.path.insert(0, str(Path(__file__).resolve().parent))
import features
from features.auto_update.announcing import announce  # noqa: E402
from commands.boot import boot  # noqa: E402
import commands.cli  # noqa: E402,F401
from commands.http import dispatch, unanswered  # noqa: E402
from features.routing import Reply  # noqa: E402
from engine import runtime  # noqa: E402
from engine.stop import asked  # noqa: E402
from engine.viewer import elsewhere, heartbeat, known, remember  # noqa: E402
from controllers.types import warm  # noqa: E402
from providers.turns import read_transcripts  # noqa: E402
from runner.chat_mirror import replay  # noqa: E402
from engine.runtime import default_env
from engine.package import ARCHIVE, CODE, ZIPPED, code_stamp, entry

DEFAULT_PORT = 8430
REQUEST_BACKLOG = 128
SWITCH_INTERVAL = 0.001
LOOPBACK = re.compile(r"^http://(?:127\.0\.0\.1|localhost)(?::(\d+))?$")


class Handler(BaseHTTPRequestHandler):
    root: Path = Path(".journal")

    def log_message(self, *_):
        pass

    def trusted_origin(self, origin: str) -> bool:
        found = LOOPBACK.match(origin)
        ports = {self.server.server_port, *(urlparse(journal.url).port for journal in known())}
        return bool(found) and int(found.group(1) or 80) in ports

    def sibling(self) -> None:
        origin = self.headers.get("Origin")
        if origin is None or not self.trusted_origin(origin):
            return
        self.send_header("Access-Control-Allow-Origin", origin)
        self.send_header("Access-Control-Allow-Methods", "GET, POST, OPTIONS")
        self.send_header("Access-Control-Allow-Headers", "Content-Type")
        self.send_header("Vary", "Origin")

    def allowed_request(self) -> bool:
        host = self.headers.get("Host")
        port = self.server.server_port
        if host not in (f"127.0.0.1:{port}", f"localhost:{port}"):
            self.send_error(403)
            return False
        origin = self.headers.get("Origin")
        if origin is not None and not self.trusted_origin(origin):
            self.send_error(403)
            return False
        return True

    def handle_one(self, method: str) -> None:
        if not self.allowed_request():
            return
        url = urlparse(self.path)
        length = self.headers["Content-Length"]
        raw = self.rfile.read(int(length)) if length else b""
        kind = self.headers.get("Content-Type") or ""
        try:
            body = {"_raw": raw, "_type": kind} if kind.startswith("multipart/") or kind.startswith("text/plain") else json.loads(raw or b"{}")
        except (json.JSONDecodeError, UnicodeDecodeError) as error:
            reply = Reply(400, {"error": f"the request body is not JSON: {error}"})
        else:
            reply = dispatch(method, url.path, self.root, dict(parse_qsl(url.query)), body)
        self.send_response(reply.code)
        self.sibling()
        self.send_header("Content-Type", reply.kind)
        if reply.chunks is None:
            data = reply.bytes()
            self.send_header("Content-Length", str(len(data)))
            self.end_headers()
            self.wfile.write(data)
            self.wfile.flush()
            if reply.after:
                reply.after()
            return
        self.send_header("Cache-Control", "no-cache")
        self.end_headers()
        if reply.after:
            reply.after()
        try:
            for chunk in reply.chunks:
                self.wfile.write(chunk)
                self.wfile.flush()
        except (BrokenPipeError, ConnectionResetError, OSError):
            reply.chunks.close()

    def do_GET(self):
        self.handle_one("GET")

    def do_POST(self):
        self.handle_one("POST")

    def do_OPTIONS(self):
        if not self.allowed_request():
            return
        self.send_response(204)
        self.sibling()
        self.end_headers()


class JournalServer(ThreadingHTTPServer):
    request_queue_size = REQUEST_BACKLOG


def serve(root: Path, port: int = DEFAULT_PORT) -> ThreadingHTTPServer:
    Handler.root = root
    other = elsewhere(root)
    if other:
        print(f"journal: this journal is already served at {other}", flush=True)
        raise SystemExit(0)
    boot(root)
    announce(root)
    features.FEATURES["plugins"].host(root)
    server = JournalServer(("127.0.0.1", port), Handler)
    remember(root, server.server_address[1])
    heartbeat(root, server.server_address[1])
    return server


WATCH_SECONDS = 1.0
SETTLE_SECONDS = 1.5
STOP_SECONDS = 0.2
FREEZE_SECONDS = 10.0
LATE_STOP = 5.0


def watch_code(root: Path, package: Path, server: ThreadingHTTPServer, changed: threading.Event) -> None:
    before = code_stamp(package)
    last_change = 0.0
    while not changed.is_set():
        time.sleep(WATCH_SECONDS)
        now = code_stamp(package)
        if now != before:
            before = now
            last_change = time.monotonic()
        elif last_change and time.monotonic() - last_change >= SETTLE_SECONDS:
            runtime.restarting(root).write_text(str(time.time()))
            changed.set()
            server.shutdown()


def watch_stop(root: Path, server: ThreadingHTTPServer, halting: threading.Event, began: float = 0.0) -> None:
    while not halting.is_set():
        time.sleep(STOP_SECONDS)
        if asked(root, began):
            halting.set()
            server.shutdown()


def watch_runtime(root: Path, halting: threading.Event) -> None:
    while not halting.wait(WATCH_SECONDS):
        runtime.refresh_flags(root)
        if runtime.hook_failures(root).is_file():
            unanswered(root)


def freeze_caches(halting: threading.Event) -> None:
    while not halting.wait(FREEZE_SECONDS):
        gc.freeze()


def warm_commands() -> None:
    from commands.cli import served
    from commands.parser import parser
    for noun in sorted(served()):
        parser(noun)


def warmed(root: Path) -> None:
    try:
        warm_viewer(root, default_env(root))
        read_transcripts(root)
    except Exception:
        traceback.print_exc()
        os._exit(1)
    gc.freeze()


def warm_viewer(root: Path, env: str) -> None:
    from commands.parser import parser
    from controllers.types import CONTROLLERS
    parser()
    dispatch("GET", f"/api/{env}/dashboard", root, {"types": ",".join(CONTROLLERS), "completed": "1", "last": "25", "events": "100"}, {})
    dispatch("GET", f"/api/{env}/family", root, {}, {})
    dispatch("GET", "/api/manifest", root, {}, {})


def run(root: Path, port: int = DEFAULT_PORT) -> None:
    sys.setswitchinterval(SWITCH_INTERVAL)
    runtime.STARTED[0] = time.time()
    server = serve(root, port)
    print(f"http://127.0.0.1:{server.server_address[1]}/", flush=True)
    changed = threading.Event()
    halting = threading.Event()
    threading.Thread(target=warmed, args=(root,), daemon=True).start()
    threading.Thread(target=watch_code, args=(root, Path(root) / ARCHIVE if ZIPPED else CODE, server, changed), daemon=True).start()
    threading.Thread(target=watch_stop, args=(root, server, halting, time.time() - LATE_STOP), daemon=True).start()
    threading.Thread(target=watch_runtime, args=(root, halting), daemon=True).start()
    threading.Thread(target=freeze_caches, args=(halting,), daemon=True).start()
    threading.Thread(target=replay, args=(root,), daemon=True).start()
    threading.Thread(target=warm, args=(root,), daemon=True).start()
    threading.Thread(target=warm_commands, daemon=True).start()
    try:
        server.serve_forever()
    except KeyboardInterrupt:
        pass
    finally:
        server.server_close()
    if halting.is_set():
        print("journal: stopped", flush=True)
        return
    if changed.is_set():
        print("journal: Python code changed; restarting on the same port", flush=True)
        command = [*entry("journal"), "--root", str(root), "serve", "--port", str(server.server_port)]
        os.execv(sys.executable, command)


if __name__ == "__main__":
    root = Path(sys.argv[1] if len(sys.argv) > 1 else ".journal").resolve()
    port = int(sys.argv[2]) if len(sys.argv) > 2 else DEFAULT_PORT
    run(root, port)
import signal
import sys
import threading
import time
from pathlib import Path

sys.path.insert(0, str(Path(__file__).resolve().parents[1]))
from providers import DRIVERS  # noqa: E402
from engine import viewer  # noqa: E402
from engine.services import Manager  # noqa: E402
from runner.engines import ENDING, supervise  # noqa: E402
from features.plugins.services import plugin_services  # noqa: E402
from engine import runtime  # noqa: E402
from controllers.faults import threw  # noqa: E402
from engine.stop import asked, session_flag  # noqa: E402
from supervisor import HEAL, RELAUNCH, RELOAD, STOP  # noqa: E402
from agents.terminal import TerminalSession, seated  # noqa: E402
from agents.actors import Agent  # noqa: E402
from engine.record import Record  # noqa: E402
from engine.package import CODE, installed_stamp, own_build  # noqa: E402
from engine.sessions import Sessions, hold_build  # noqa: E402
import features  # noqa: E402
from features.switches import watch_change_log  # noqa: E402
from features.auto_update.check import UpdateCheck  # noqa: E402
from features.auto_update.relaunch import Relaunch  # noqa: E402
from features.work_tracking.auto import CheckIn  # noqa: E402

TICK = 0.25
RELOAD_EVERY = 1.0
VIEWER_EVERY = 10.0
SERVICES_EVERY = 1.0
CHECKS_EVERY = 1.0
SERVER_CRASHES = 3
RETRY_AFTER = 1.0
STARTUP, EARLY = 30.0, 16384
CONSENT_EVERY = 3.0
TERMINATED = threading.Event()


def keep_viewer(root: Path, cwd: Path, watching, exits: list) -> object:
    if watching and watching.is_alive():
        return watching
    thread = threading.Thread(target=lambda: exits.append(viewer.launch(root, cwd)[1]), daemon=True)
    thread.start()
    return thread


def moved(seat: TerminalSession) -> bool:
    return Sessions(seat.root).environment(seat.session) not in ("", seat.env)


def crashing(exits: list) -> bool:
    return len(exits) >= SERVER_CRASHES and all(exits[-SERVER_CRASHES:])


def checks(seat: TerminalSession, driver) -> list:
    try:
        features.load()
        watcher = Agent(driver.record, driver)
        return [CheckIn(watcher), UpdateCheck(watcher), Relaunch(watcher)]
    except Exception:
        threw(seat.root, seat.env, "starting the worker's checks")
        return []


def run_checks(seat: TerminalSession, driver, kept: list) -> None:
    for step in [driver.pump, *(check.tick for check in kept)] if kept else []:
        try:
            step()
        except Exception:
            threw(seat.root, seat.env, f"a worker check: {type(getattr(step, '__self__', step)).__name__}")


class Confirm:
    def __init__(self, driver):
        self.driver = driver
        self.at = self.driver.printed.stat().st_size if self.driver.printed.is_file() else 0
        self.started = self.consented = time.time()
        self.ready = 0.0
        self.answered = False

    def tick(self) -> None:
        if self.answered or time.time() - self.started >= STARTUP or not self.driver.printed.is_file():
            return
        fresh = self.driver.printed.stat().st_size - self.at
        early = self.driver.printed_tail(min(fresh, EARLY)) if fresh > 0 else b""
        if self.driver.consent(early):
            self.consent(early)
            return
        opening = self.driver.opening(early)
        self.ready = (self.ready or time.time()) if opening else 0.0
        if opening and time.time() - self.ready >= self.driver.CONFIRM_AFTER:
            self.answered = True
            self.driver.send(opening, now=True)

    def consent(self, early: bytes) -> None:
        if time.time() - self.consented < CONSENT_EVERY:
            return
        self.consented = time.time()
        self.driver.press_raw(self.driver.consent(early))


def run(root: Path, cwd: Path, env: str, agent: str, session: str, lifeline: int = -1) -> int:
    hold_build(root, CODE)
    watch_change_log()
    seat = seated(TerminalSession(root, env, agent, session))
    relaunching = runtime.relaunch_file(root, session)
    stopping = session_flag(root, session)
    stamps = installed_stamp(root)
    began = time.time()
    driver = DRIVERS[agent](Record(root, env), session)
    confirm = Confirm(driver)
    kept = checks(seat, driver)
    services = Manager(root, lifeline, sources=(plugin_services,))
    last_check = last_viewer = last_services = last_checks = 0.0
    watching = None
    exits: list = []
    keeps = own_build(root)
    engines = threading.Event()
    supervising = threading.Thread(target=supervise, args=(root, engines), daemon=True)
    if keeps:
        supervising.start()
    try:
        while True:
            time.sleep(TICK)
            confirm.tick()
            if asked(root, began) or stopping.is_file() or TERMINATED.is_set():
                stopping.unlink(missing_ok=True)
                return STOP
            if relaunching.is_file():
                return RELAUNCH
            now = time.time()
            if keeps and now - last_services >= SERVICES_EVERY:
                last_services = now
                services.tick()
            if keeps and now - last_viewer >= (RETRY_AFTER if exits and exits[-1] else VIEWER_EVERY):
                last_viewer = now
                watching = keep_viewer(root, cwd, watching, exits)
                if crashing(exits):
                    return HEAL
            if now - last_checks >= CHECKS_EVERY:
                last_checks = now
                if moved(seat):
                    return RELOAD
                if not kept:
                    kept = checks(seat, driver)
                run_checks(seat, driver, kept)
            if now - last_check >= RELOAD_EVERY:
                last_check = now
                if installed_stamp(root) != stamps and not runtime.upgrading(root):
                    return RELOAD
    finally:
        engines.set()
        if keeps:
            supervising.join(timeout=ENDING)


def ended(signum, frame) -> None:
    TERMINATED.set()


if __name__ == "__main__":
    signal.signal(signal.SIGTERM, ended)
    root, cwd, env, agent, session = sys.argv[1:6]
    raise SystemExit(run(Path(root), Path(cwd), env, agent, session, int(sys.argv[6]) if len(sys.argv) > 6 else -1))
