| 123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185 |
- from __future__ import annotations
- import asyncio
- import os
- from collections.abc import Mapping
- from typing import Any
- import httpx
- _TERMINAL_TRACE_STATUSES = {"completed", "failed", "stopped", "cancelled"}
- class HostApiError(RuntimeError):
- def __init__(self, status_code: int, message: str) -> None:
- super().__init__(message)
- self.status_code = status_code
- class HostClient:
- def __init__(
- self,
- base_url: str | None = None,
- *,
- transport: httpx.AsyncBaseTransport | None = None,
- ) -> None:
- resolved = base_url or os.getenv("SCRIPT_BUILD_HOST_API_BASE") or "http://127.0.0.1:8080"
- self._client = httpx.AsyncClient(
- base_url=resolved.rstrip("/"),
- timeout=20.0,
- transport=transport,
- )
- self._inputs: dict[int, dict[str, Any]] = {}
- self._contracts_by_spec: dict[tuple[int, str, int], dict[str, Any]] = {}
- self._events: dict[int, list[dict[str, Any]]] = {}
- self._event_cursors: dict[int, str] = {}
- self._messages: dict[tuple[int, str], list[dict[str, Any]]] = {}
- async def aclose(self) -> None:
- await self._client.aclose()
- async def list_runs(self, headers: Mapping[str, str]) -> dict[str, Any]:
- return await self._get("/api/pattern/script_builds?page=1&page_size=100", headers)
- async def journey_source(
- self, script_build_id: int, headers: Mapping[str, str]
- ) -> dict[str, Any]:
- prefix = f"/api/pattern/script_builds/{script_build_id}"
- snapshot, traces = await asyncio.gather(
- self._get(f"{prefix}/mission", headers),
- self._get(f"{prefix}/traces", headers),
- )
- input_summary, events = await asyncio.gather(
- self._input_summary(script_build_id, prefix, headers),
- self._incremental_events(script_build_id, f"{prefix}/mission/events", headers),
- )
- tasks = snapshot.get("tasks") or []
- contract_tasks = [
- task
- for task in tasks
- if any(
- str(ref).startswith("script-build://task-contracts/sha256/")
- for ref in (task.get("current_spec") or {}).get("context_refs") or []
- )
- ]
- trace_values = traces.get("traces") or []
- contracts, messages = await asyncio.gather(
- self._contracts(script_build_id, prefix, contract_tasks, headers),
- self._trace_messages(script_build_id, prefix, trace_values, headers),
- )
- return {
- "snapshot": snapshot,
- "input_summary": input_summary,
- "events": events,
- "messages": messages,
- "contracts": contracts,
- }
- async def artifact(
- self, script_build_id: int, artifact_version_id: int, headers: Mapping[str, str]
- ) -> dict[str, Any]:
- return await self._get(
- f"/api/pattern/script_builds/{script_build_id}/artifacts/{artifact_version_id}",
- headers,
- )
- async def _input_summary(
- self, script_build_id: int, prefix: str, headers: Mapping[str, str]
- ) -> dict[str, Any]:
- cached = self._inputs.get(script_build_id)
- if cached is None:
- cached = await self._get(f"{prefix}/input-summary", headers)
- self._inputs[script_build_id] = cached
- return cached
- async def _contracts(
- self,
- script_build_id: int,
- prefix: str,
- tasks: list[dict[str, Any]],
- headers: Mapping[str, str],
- ) -> dict[str, dict[str, Any]]:
- async def read(task: dict[str, Any]) -> dict[str, Any]:
- task_id = str(task["task_id"])
- version = int((task.get("current_spec") or {}).get("version") or 0)
- key = (script_build_id, task_id, version)
- cached = self._contracts_by_spec.get(key)
- if cached is None:
- cached = await self._get(f"{prefix}/tasks/{task_id}/contract", headers)
- self._contracts_by_spec[key] = cached
- return cached
- values = await asyncio.gather(*(read(task) for task in tasks))
- return {str(value["task_id"]): value for value in values}
- async def _trace_messages(
- self,
- script_build_id: int,
- prefix: str,
- traces: list[dict[str, Any]],
- headers: Mapping[str, str],
- ) -> dict[str, list[dict[str, Any]]]:
- async def read(trace: dict[str, Any]) -> tuple[str, list[dict[str, Any]]]:
- trace_id = str(trace["trace_id"])
- key = (script_build_id, trace_id)
- cached = self._messages.get(key, [])
- status = str(trace.get("status") or "").lower()
- if cached and status in _TERMINAL_TRACE_STATUSES:
- return trace_id, cached
- after = max((int(item.get("sequence") or 0) for item in cached), default=0)
- page = await self._get(
- f"{prefix}/traces/{trace_id}/messages?after_sequence={after}", headers
- )
- merged = _merge_by_sequence(cached, list(page.get("messages") or []))
- self._messages[key] = merged
- return trace_id, merged
- values = await asyncio.gather(*(read(trace) for trace in traces if trace.get("trace_id")))
- return dict(values)
- async def _incremental_events(
- self, script_build_id: int, path: str, headers: Mapping[str, str]
- ) -> list[dict[str, Any]]:
- events = list(self._events.get(script_build_id, []))
- cursor = self._event_cursors.get(script_build_id)
- for _ in range(100):
- query = "?limit=500"
- if cursor:
- query += f"&after={cursor}"
- page = await self._get(path + query, headers)
- additions = list(page.get("events") or [])
- events.extend(additions)
- next_cursor = page.get("next_cursor")
- if next_cursor:
- cursor = str(next_cursor)
- self._event_cursors[script_build_id] = cursor
- if len(additions) < 500:
- break
- self._events[script_build_id] = events
- return events
- async def _get(self, path: str, headers: Mapping[str, str]) -> dict[str, Any]:
- forwarded = {
- key: value
- for key, value in headers.items()
- if key.lower() in {"authorization", "cookie", "x-request-id"}
- }
- response = await self._client.get(path, headers=forwarded)
- if not response.is_success:
- try:
- payload = response.json()
- message = str(payload.get("message") or payload.get("detail") or response.text)
- except ValueError:
- message = response.text
- raise HostApiError(response.status_code, message or "Host request failed")
- value = response.json()
- if not isinstance(value, dict):
- raise HostApiError(502, "Host returned a non-object response")
- return value
- def _merge_by_sequence(
- current: list[dict[str, Any]], additions: list[dict[str, Any]]
- ) -> list[dict[str, Any]]:
- merged = {int(item.get("sequence") or 0): item for item in [*current, *additions]}
- return [merged[key] for key in sorted(merged)]
|