diff options
Diffstat (limited to 'src/gap_pipeline/clients.py')
| -rw-r--r-- | src/gap_pipeline/clients.py | 168 |
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]}", + ) |
