host_client.py 7.1 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185
  1. from __future__ import annotations
  2. import asyncio
  3. import os
  4. from collections.abc import Mapping
  5. from typing import Any
  6. import httpx
  7. _TERMINAL_TRACE_STATUSES = {"completed", "failed", "stopped", "cancelled"}
  8. class HostApiError(RuntimeError):
  9. def __init__(self, status_code: int, message: str) -> None:
  10. super().__init__(message)
  11. self.status_code = status_code
  12. class HostClient:
  13. def __init__(
  14. self,
  15. base_url: str | None = None,
  16. *,
  17. transport: httpx.AsyncBaseTransport | None = None,
  18. ) -> None:
  19. resolved = base_url or os.getenv("SCRIPT_BUILD_HOST_API_BASE") or "http://127.0.0.1:8080"
  20. self._client = httpx.AsyncClient(
  21. base_url=resolved.rstrip("/"),
  22. timeout=20.0,
  23. transport=transport,
  24. )
  25. self._inputs: dict[int, dict[str, Any]] = {}
  26. self._contracts_by_spec: dict[tuple[int, str, int], dict[str, Any]] = {}
  27. self._events: dict[int, list[dict[str, Any]]] = {}
  28. self._event_cursors: dict[int, str] = {}
  29. self._messages: dict[tuple[int, str], list[dict[str, Any]]] = {}
  30. async def aclose(self) -> None:
  31. await self._client.aclose()
  32. async def list_runs(self, headers: Mapping[str, str]) -> dict[str, Any]:
  33. return await self._get("/api/pattern/script_builds?page=1&page_size=100", headers)
  34. async def journey_source(
  35. self, script_build_id: int, headers: Mapping[str, str]
  36. ) -> dict[str, Any]:
  37. prefix = f"/api/pattern/script_builds/{script_build_id}"
  38. snapshot, traces = await asyncio.gather(
  39. self._get(f"{prefix}/mission", headers),
  40. self._get(f"{prefix}/traces", headers),
  41. )
  42. input_summary, events = await asyncio.gather(
  43. self._input_summary(script_build_id, prefix, headers),
  44. self._incremental_events(script_build_id, f"{prefix}/mission/events", headers),
  45. )
  46. tasks = snapshot.get("tasks") or []
  47. contract_tasks = [
  48. task
  49. for task in tasks
  50. if any(
  51. str(ref).startswith("script-build://task-contracts/sha256/")
  52. for ref in (task.get("current_spec") or {}).get("context_refs") or []
  53. )
  54. ]
  55. trace_values = traces.get("traces") or []
  56. contracts, messages = await asyncio.gather(
  57. self._contracts(script_build_id, prefix, contract_tasks, headers),
  58. self._trace_messages(script_build_id, prefix, trace_values, headers),
  59. )
  60. return {
  61. "snapshot": snapshot,
  62. "input_summary": input_summary,
  63. "events": events,
  64. "messages": messages,
  65. "contracts": contracts,
  66. }
  67. async def artifact(
  68. self, script_build_id: int, artifact_version_id: int, headers: Mapping[str, str]
  69. ) -> dict[str, Any]:
  70. return await self._get(
  71. f"/api/pattern/script_builds/{script_build_id}/artifacts/{artifact_version_id}",
  72. headers,
  73. )
  74. async def _input_summary(
  75. self, script_build_id: int, prefix: str, headers: Mapping[str, str]
  76. ) -> dict[str, Any]:
  77. cached = self._inputs.get(script_build_id)
  78. if cached is None:
  79. cached = await self._get(f"{prefix}/input-summary", headers)
  80. self._inputs[script_build_id] = cached
  81. return cached
  82. async def _contracts(
  83. self,
  84. script_build_id: int,
  85. prefix: str,
  86. tasks: list[dict[str, Any]],
  87. headers: Mapping[str, str],
  88. ) -> dict[str, dict[str, Any]]:
  89. async def read(task: dict[str, Any]) -> dict[str, Any]:
  90. task_id = str(task["task_id"])
  91. version = int((task.get("current_spec") or {}).get("version") or 0)
  92. key = (script_build_id, task_id, version)
  93. cached = self._contracts_by_spec.get(key)
  94. if cached is None:
  95. cached = await self._get(f"{prefix}/tasks/{task_id}/contract", headers)
  96. self._contracts_by_spec[key] = cached
  97. return cached
  98. values = await asyncio.gather(*(read(task) for task in tasks))
  99. return {str(value["task_id"]): value for value in values}
  100. async def _trace_messages(
  101. self,
  102. script_build_id: int,
  103. prefix: str,
  104. traces: list[dict[str, Any]],
  105. headers: Mapping[str, str],
  106. ) -> dict[str, list[dict[str, Any]]]:
  107. async def read(trace: dict[str, Any]) -> tuple[str, list[dict[str, Any]]]:
  108. trace_id = str(trace["trace_id"])
  109. key = (script_build_id, trace_id)
  110. cached = self._messages.get(key, [])
  111. status = str(trace.get("status") or "").lower()
  112. if cached and status in _TERMINAL_TRACE_STATUSES:
  113. return trace_id, cached
  114. after = max((int(item.get("sequence") or 0) for item in cached), default=0)
  115. page = await self._get(
  116. f"{prefix}/traces/{trace_id}/messages?after_sequence={after}", headers
  117. )
  118. merged = _merge_by_sequence(cached, list(page.get("messages") or []))
  119. self._messages[key] = merged
  120. return trace_id, merged
  121. values = await asyncio.gather(*(read(trace) for trace in traces if trace.get("trace_id")))
  122. return dict(values)
  123. async def _incremental_events(
  124. self, script_build_id: int, path: str, headers: Mapping[str, str]
  125. ) -> list[dict[str, Any]]:
  126. events = list(self._events.get(script_build_id, []))
  127. cursor = self._event_cursors.get(script_build_id)
  128. for _ in range(100):
  129. query = "?limit=500"
  130. if cursor:
  131. query += f"&after={cursor}"
  132. page = await self._get(path + query, headers)
  133. additions = list(page.get("events") or [])
  134. events.extend(additions)
  135. next_cursor = page.get("next_cursor")
  136. if next_cursor:
  137. cursor = str(next_cursor)
  138. self._event_cursors[script_build_id] = cursor
  139. if len(additions) < 500:
  140. break
  141. self._events[script_build_id] = events
  142. return events
  143. async def _get(self, path: str, headers: Mapping[str, str]) -> dict[str, Any]:
  144. forwarded = {
  145. key: value
  146. for key, value in headers.items()
  147. if key.lower() in {"authorization", "cookie", "x-request-id"}
  148. }
  149. response = await self._client.get(path, headers=forwarded)
  150. if not response.is_success:
  151. try:
  152. payload = response.json()
  153. message = str(payload.get("message") or payload.get("detail") or response.text)
  154. except ValueError:
  155. message = response.text
  156. raise HostApiError(response.status_code, message or "Host request failed")
  157. value = response.json()
  158. if not isinstance(value, dict):
  159. raise HostApiError(502, "Host returned a non-object response")
  160. return value
  161. def _merge_by_sequence(
  162. current: list[dict[str, Any]], additions: list[dict[str, Any]]
  163. ) -> list[dict[str, Any]]:
  164. merged = {int(item.get("sequence") or 0): item for item in [*current, *additions]}
  165. return [merged[key] for key in sorted(merged)]