summaryrefslogtreecommitdiff
path: root/src/gap_pipeline/clients.py
diff options
context:
space:
mode:
authorAnonymous Authors <anonymous@invalid.example>2026-07-24 13:24:36 -0500
committerAnonymous Authors <anonymous@invalid.example>2026-07-24 13:24:36 -0500
commitdb293f3606a97b3e417de27124858e134005acbd (patch)
tree8efeedcd2033b82d1c90eb0cb84e134421ff1a8f /src/gap_pipeline/clients.py
Add minimal GAP reproduction package
Diffstat (limited to 'src/gap_pipeline/clients.py')
-rw-r--r--src/gap_pipeline/clients.py168
1 files changed, 168 insertions, 0 deletions
diff --git a/src/gap_pipeline/clients.py b/src/gap_pipeline/clients.py
new file mode 100644
index 0000000..3b3e40f
--- /dev/null
+++ b/src/gap_pipeline/clients.py
@@ -0,0 +1,168 @@
+"""JSON LLM adapters. Offline tests use ReplayClient or ScriptedClient."""
+
+from __future__ import annotations
+
+import json
+from collections import defaultdict, deque
+from dataclasses import dataclass, field
+from pathlib import Path
+from typing import Any, Protocol
+
+
+@dataclass(frozen=True)
+class ModelResponse:
+ data: dict[str, Any]
+ raw_text: str
+ model: str
+ response_id: str | None = None
+ usage: dict[str, Any] = field(default_factory=dict)
+
+
+class JsonLLM(Protocol):
+ model: str
+
+ async def generate_json(
+ self,
+ *,
+ system_prompt: str,
+ user_prompt: str,
+ request_id: str,
+ ) -> ModelResponse: ...
+
+
+class OpenAIJsonClient:
+ """Optional OpenAI chat-completions adapter.
+
+ No API request is made until ``generate_json`` is awaited.
+ """
+
+ def __init__(
+ self,
+ model: str,
+ *,
+ api_key: str | None = None,
+ max_completion_tokens: int = 16_000,
+ seed: int | None = None,
+ ) -> None:
+ try:
+ from openai import AsyncOpenAI
+ except ImportError as exc:
+ raise RuntimeError(
+ "Install the optional API dependency: pip install -e '.[api]'"
+ ) from exc
+ self.model = model
+ self._client = AsyncOpenAI(api_key=api_key)
+ self.max_completion_tokens = max_completion_tokens
+ self.seed = seed
+
+ async def generate_json(
+ self,
+ *,
+ system_prompt: str,
+ user_prompt: str,
+ request_id: str,
+ ) -> ModelResponse:
+ kwargs: dict[str, Any] = {
+ "model": self.model,
+ "messages": [
+ {"role": "system", "content": system_prompt},
+ {"role": "user", "content": user_prompt},
+ ],
+ "response_format": {"type": "json_object"},
+ "max_completion_tokens": self.max_completion_tokens,
+ }
+ if self.seed is not None:
+ kwargs["seed"] = self.seed
+ response = await self._client.chat.completions.create(**kwargs)
+ raw = response.choices[0].message.content or "{}"
+ data = json.loads(raw)
+ usage = (
+ response.usage.model_dump()
+ if response.usage is not None and hasattr(response.usage, "model_dump")
+ else {}
+ )
+ return ModelResponse(
+ data=data,
+ raw_text=raw,
+ model=self.model,
+ response_id=response.id,
+ usage=usage,
+ )
+
+
+class ReplayClient:
+ """Replay archived call responses keyed by request ID."""
+
+ def __init__(self, path: Path, model: str = "replay") -> None:
+ self.model = model
+ payload = json.loads(path.read_text(encoding="utf-8"))
+ if isinstance(payload, list):
+ self._responses = {
+ str(row["request_id"]): row["response_data"] for row in payload
+ }
+ elif isinstance(payload, dict):
+ self._responses = payload
+ else:
+ raise ValueError("replay file must be a JSON object or list")
+
+ async def generate_json(
+ self,
+ *,
+ system_prompt: str,
+ user_prompt: str,
+ request_id: str,
+ ) -> ModelResponse:
+ if request_id not in self._responses:
+ raise KeyError(f"no replay response for {request_id}")
+ data = self._responses[request_id]
+ return ModelResponse(
+ data=data,
+ raw_text=json.dumps(data, ensure_ascii=False),
+ model=self.model,
+ response_id=f"replay:{request_id}",
+ )
+
+
+class ScriptedClient:
+ """Small deterministic client for tests and offline smoke runs."""
+
+ def __init__(
+ self,
+ responses: dict[str, list[dict[str, Any]] | dict[str, Any]],
+ model: str = "scripted",
+ ) -> None:
+ self.model = model
+ self.calls: list[str] = []
+ self._responses: dict[str, deque[dict[str, Any]]] = {}
+ for prefix, values in responses.items():
+ sequence = values if isinstance(values, list) else [values]
+ self._responses[prefix] = deque(sequence)
+ self._prefix_counts: defaultdict[str, int] = defaultdict(int)
+
+ async def generate_json(
+ self,
+ *,
+ system_prompt: str,
+ user_prompt: str,
+ request_id: str,
+ ) -> ModelResponse:
+ self.calls.append(request_id)
+ prefixes = sorted(
+ (prefix for prefix in self._responses if request_id.startswith(prefix)),
+ key=len,
+ reverse=True,
+ )
+ if not prefixes:
+ raise KeyError(f"no scripted response matches {request_id}")
+ prefix = prefixes[0]
+ queue = self._responses[prefix]
+ if not queue:
+ raise IndexError(f"scripted responses exhausted for {prefix}")
+ data = queue.popleft()
+ self._prefix_counts[prefix] += 1
+ return ModelResponse(
+ data=data,
+ raw_text=json.dumps(data, ensure_ascii=False),
+ model=self.model,
+ response_id=f"scripted:{prefix}:{self._prefix_counts[prefix]}",
+ )