diff --git a/ai_research/dataset_generation/cli.py b/ai_research/dataset_generation/cli.py index 860ca4e..a0ecca2 100644 --- a/ai_research/dataset_generation/cli.py +++ b/ai_research/dataset_generation/cli.py @@ -7,6 +7,9 @@ from ai_research.dataset_generation.domain.category import get_random_category from ai_research.dataset_generation.domain.llm_response import parse_examples +from ai_research.dataset_generation.infrastructure.openai_codex_client import ( + OpenAICodexTextClient, +) from ai_research.dataset_generation.infrastructure.openai_client import OpenAITextClient from ai_research.dataset_generation.infrastructure.prompts import ( get_filtering_prompt, @@ -18,18 +21,40 @@ logger = logging.getLogger(__name__) +LOG_FORMAT = "%(asctime)s %(levelname)s %(name)s: %(message)s" +LOG_DATE_FORMAT = "%Y-%m-%dT%H:%M:%S%z" +PROMPT_ARTIFACT_MARKERS = ( + "[Assertion]", + "[Code]", + "[Thinking]", + "[Explanation]", + "Example Set", +) +PROMPT_ARTIFACT_MARKERS_LOWER = tuple( + marker.lower() for marker in PROMPT_ARTIFACT_MARKERS +) + def _add_llm_arguments(parser: argparse.ArgumentParser) -> None: parser.add_argument("--model", default="gpt-4o-mini") parser.add_argument("--max-tokens", type=int, default=32768) parser.add_argument("--temperature", type=float, default=1.0) - parser.add_argument("--api-key", required=True) + parser.add_argument("--api-key", default="") parser.add_argument("--base-url") -def _build_client(args: argparse.Namespace) -> OpenAITextClient: - return OpenAITextClient( - api_key=args.api_key, +def _build_client(args: argparse.Namespace) -> OpenAITextClient | OpenAICodexTextClient: + if args.api_key: + return OpenAITextClient( + api_key=args.api_key, + base_url=args.base_url, + model=args.model, + temperature=args.temperature, + max_tokens=args.max_tokens, + ) + + return OpenAICodexTextClient( + api_key="", base_url=args.base_url, model=args.model, temperature=args.temperature, @@ -37,9 +62,23 @@ def _build_client(args: argparse.Namespace) -> OpenAITextClient: ) +def _contains_prompt_artifact(value: object) -> bool: + if isinstance(value, str): + lowered = value.lower() + return any(marker in lowered for marker in PROMPT_ARTIFACT_MARKERS_LOWER) + if isinstance(value, dict): + return any(_contains_prompt_artifact(item) for item in value.values()) + if isinstance(value, list): + return any(_contains_prompt_artifact(item) for item in value) + return False + + +TextClient = OpenAITextClient | OpenAICodexTextClient + + def run_generate( args: argparse.Namespace, - client_factory: Callable[[argparse.Namespace], OpenAITextClient] = _build_client, + client_factory: Callable[[argparse.Namespace], TextClient] = _build_client, category_picker: Callable[[], str] = get_random_category, ) -> int: template = get_generation_prompt() @@ -59,7 +98,7 @@ def run_generate( def run_filter( args: argparse.Namespace, - client_factory: Callable[[argparse.Namespace], OpenAITextClient] = _build_client, + client_factory: Callable[[argparse.Namespace], TextClient] = _build_client, ) -> int: template = get_filtering_prompt() client = client_factory(args) @@ -84,13 +123,21 @@ def run_filter( def run_transform(args: argparse.Namespace) -> int: with Path(args.input).open() as in_file, Path(args.output).open("w") as out_file: - for raw_line in in_file: + for line_number, raw_line in enumerate(in_file, start=1): line = raw_line.strip() if not line: continue output_line = json.loads(line) - for example in output_line.get("examples", []): + for example_index, example in enumerate(output_line.get("examples", []), start=1): + if _contains_prompt_artifact(example): + logger.warning( + "Rejected mis-parsed generated example at line %s, example %s", + line_number, + example_index, + ) + continue + json.dump(example, out_file) out_file.write("\n") @@ -131,5 +178,10 @@ def build_parser() -> argparse.ArgumentParser: def main(argv: Optional[Sequence[str]] = None) -> int: parser = build_parser() args = parser.parse_args(argv) - logging.basicConfig(level=logging.INFO) + logging.basicConfig( + level=logging.INFO, + format=LOG_FORMAT, + datefmt=LOG_DATE_FORMAT, + ) + logger.info("Starting dataset generation with command: %s", args.command) return args.func(args) diff --git a/ai_research/dataset_generation/infrastructure/openai_client.py b/ai_research/dataset_generation/infrastructure/openai_client.py index 86d7c2d..7bcc4f0 100644 --- a/ai_research/dataset_generation/infrastructure/openai_client.py +++ b/ai_research/dataset_generation/infrastructure/openai_client.py @@ -64,27 +64,16 @@ def __init__( def infer(self, template: str, data: Mapping[str, object]) -> str: prompt = render_prompt_template(template, data) delay = self.retry_delay_seconds - use_max_completion_tokens = False for attempt in range(self.max_retries + 1): try: messages: list[ChatCompletionUserMessageParam] = [ {"role": "user", "content": prompt} ] - if use_max_completion_tokens: - response = self._client.chat.completions.create( - model=self.model, - messages=messages, - temperature=self.temperature, - max_completion_tokens=self.max_tokens, - ) - else: - response = self._client.chat.completions.create( - model=self.model, - messages=messages, - temperature=self.temperature, - max_tokens=self.max_tokens, - ) + response = self._client.chat.completions.create( + model=self.model, + messages=messages, + ) return _extract_response_text(response) except ( RateLimitError, @@ -95,13 +84,6 @@ def infer(self, template: str, data: Mapping[str, object]) -> str: if attempt == self.max_retries: raise except APIStatusError as exc: - if ( - not use_max_completion_tokens - and _should_retry_with_max_completion_tokens(exc) - ): - use_max_completion_tokens = True - continue - status_code = exc.status_code or 0 if status_code not in {408, 409, 429} and status_code < 500: raise diff --git a/ai_research/dataset_generation/infrastructure/openai_codex_client.py b/ai_research/dataset_generation/infrastructure/openai_codex_client.py new file mode 100644 index 0000000..aa62c2b --- /dev/null +++ b/ai_research/dataset_generation/infrastructure/openai_codex_client.py @@ -0,0 +1,426 @@ +import json +import os +import time +import uuid +from collections.abc import Callable, Iterable, Mapping +from dataclasses import dataclass +from datetime import datetime, timezone +from pathlib import Path +from typing import Any, Optional +from urllib.error import HTTPError, URLError +from urllib.request import Request, urlopen + + +CODEX_BASE_URL = "https://chatgpt.com/backend-api/codex" +REFRESH_URL = "https://auth.openai.com/oauth/token" +REFRESH_CLIENT_ID = "app_EMoamEEZ73f0CkXaXp7hrann" +STALE_TOKEN_SECONDS = 8 * 24 * 60 * 60 + + +def render_prompt_template(template: str, data: Mapping[str, object]) -> str: + rendered = template + for key, value in data.items(): + rendered = rendered.replace(f"{{{{.{key}}}}}", str(value)) + return rendered + + +@dataclass +class _Credentials: + access_token: str + refresh_token: Optional[str] + account_id: Optional[str] + last_refresh: Optional[datetime] + + +class _HttpStatusError(RuntimeError): + def __init__(self, status_code: int, message: str) -> None: + super().__init__(message) + self.status_code = status_code + + +class _RetryableCodexStreamError(RuntimeError): + pass + + +class OpenAICodexTextClient: + def __init__( + self, + api_key: str, + model: str, + temperature: float, + max_tokens: int, + base_url: Optional[str] = None, + max_retries: int = 6, + retry_delay_seconds: float = 1.0, + http_timeout: float = 60.0, + ) -> None: + if api_key != "": + raise ValueError( + "OpenAICodexTextClient requires empty api_key; use OpenAITextClient for API-key mode" + ) + + self.api_key = api_key + self.base_url = base_url + self.model = model + self.temperature = temperature + self.max_tokens = max_tokens + self.max_retries = max_retries + self.retry_delay_seconds = retry_delay_seconds + self.http_timeout = http_timeout + self.conversation_id = str(uuid.uuid4()) + self._sleep: Callable[[float], None] = time.sleep + + def infer(self, template: str, data: Mapping[str, object]) -> str: + prompt = render_prompt_template(template, data) + credentials = self._load_credentials() + if self._is_stale(credentials): + credentials = self._refresh_credentials(credentials) + + delay = self.retry_delay_seconds + refreshed_after_401 = False + + for attempt in range(self.max_retries + 1): + try: + body = self._build_request_body(prompt) + headers = self._build_headers(credentials) + stream = self._post_stream(self._responses_url(), headers, body) + try: + return self._parse_sse(stream) + finally: + close = getattr(stream, "close", None) + if callable(close): + close() + except _HttpStatusError as exc: + if exc.status_code == 401 and not refreshed_after_401: + credentials = self._refresh_credentials(credentials) + refreshed_after_401 = True + continue + if not self._is_retryable_status(exc.status_code): + raise RuntimeError( + f"Codex inference returned non-retryable HTTP status {exc.status_code}" + ) from exc + if attempt == self.max_retries: + raise RuntimeError("Codex inference exhausted retries") from exc + except (TimeoutError, URLError, OSError) as exc: + if attempt == self.max_retries: + raise RuntimeError("Codex inference exhausted retries") from exc + except _RetryableCodexStreamError as exc: + if attempt == self.max_retries: + raise RuntimeError("Codex inference exhausted retries") from exc + + self._sleep(delay) + delay *= 2 + + raise RuntimeError("Codex inference exhausted retries") + + def _codex_home(self) -> Path: + return Path(os.environ.get("CODEX_HOME", "~/.codex")).expanduser() + + def _auth_path(self) -> Path: + return self._codex_home() / "auth.json" + + def _load_auth_payload(self) -> dict[str, Any]: + auth_path = self._auth_path() + try: + with auth_path.open() as auth_file: + payload = json.load(auth_file) + except FileNotFoundError as exc: + raise RuntimeError( + "Missing file-backed Codex auth. Use Codex ChatGPT file auth or pass --api-key for OpenAI mode." + ) from exc + except (OSError, json.JSONDecodeError) as exc: + raise RuntimeError("Malformed Codex auth.json") from exc + + if not isinstance(payload, dict): + raise RuntimeError("Malformed Codex auth.json") + return payload + + def _load_credentials(self) -> _Credentials: + payload = self._load_auth_payload() + tokens = payload.get("tokens") + auth_mode = payload.get("auth_mode") + api_key = payload.get("OPENAI_API_KEY") + if ( + auth_mode not in {None, "chatgpt"} + or api_key + or not isinstance(tokens, dict) + ): + raise RuntimeError("Malformed Codex auth.json") + + access_token = tokens.get("access_token") + if not isinstance(access_token, str) or not access_token: + raise RuntimeError("Missing tokens.access_token in Codex ChatGPT auth") + + refresh_token = tokens.get("refresh_token") + if refresh_token is not None and not isinstance(refresh_token, str): + raise RuntimeError("Malformed Codex auth.json") + + account_id = tokens.get("account_id") + if account_id is not None and not isinstance(account_id, str): + raise RuntimeError("Malformed Codex auth.json") + + last_refresh = payload.get("last_refresh") + parsed_last_refresh = None + if last_refresh is not None: + if not isinstance(last_refresh, str): + raise RuntimeError("Malformed Codex auth.json") + parsed_last_refresh = self._parse_timestamp(last_refresh) + + return _Credentials( + access_token=access_token, + refresh_token=refresh_token, + account_id=account_id, + last_refresh=parsed_last_refresh, + ) + + def _parse_timestamp(self, value: str) -> datetime: + try: + parsed = datetime.fromisoformat(value.replace("Z", "+00:00")) + except ValueError as exc: + raise RuntimeError("Malformed Codex auth.json") from exc + if parsed.tzinfo is None: + return parsed.replace(tzinfo=timezone.utc) + return parsed.astimezone(timezone.utc) + + def _is_stale(self, credentials: _Credentials) -> bool: + if credentials.last_refresh is None: + return False + age = datetime.now(timezone.utc) - credentials.last_refresh + return age.total_seconds() > STALE_TOKEN_SECONDS + + def _refresh_credentials(self, credentials: _Credentials) -> _Credentials: + if not credentials.refresh_token: + raise RuntimeError( + "Missing tokens.refresh_token required to refresh Codex ChatGPT auth" + ) + + request_body = { + "client_id": REFRESH_CLIENT_ID, + "grant_type": "refresh_token", + "refresh_token": credentials.refresh_token, + "scope": "openid profile email", + } + try: + response = self._post_json(REFRESH_URL, request_body) + except _HttpStatusError as exc: + raise RuntimeError( + f"Codex ChatGPT auth refresh returned HTTP status {exc.status_code}" + ) from exc + except (TimeoutError, URLError, OSError) as exc: + raise RuntimeError("Codex ChatGPT auth refresh failed") from exc + + if not isinstance(response, dict): + raise RuntimeError("Codex ChatGPT auth refresh returned malformed response") + + new_access_token = response.get("access_token") + if not isinstance(new_access_token, str) or not new_access_token: + raise RuntimeError("Codex ChatGPT auth refresh returned no access token") + + auth_payload = self._load_auth_payload() + tokens = auth_payload.get("tokens") + if not isinstance(tokens, dict): + raise RuntimeError("Malformed Codex auth.json") + + for field in ("id_token", "access_token", "refresh_token"): + value = response.get(field) + if isinstance(value, str) and value: + tokens[field] = value + + account_id = response.get("account_id") + if isinstance(account_id, str) and account_id: + tokens["account_id"] = account_id + auth_payload["last_refresh"] = datetime.now(timezone.utc).isoformat() + self._write_auth_payload(auth_payload) + return self._load_credentials() + + def _write_auth_payload(self, payload: Mapping[str, Any]) -> None: + auth_path = self._auth_path() + try: + auth_path.write_text(json.dumps(payload, indent=2) + "\n") + os.chmod(auth_path, 0o600) + except OSError as exc: + raise RuntimeError("Failed to persist Codex auth refresh") from exc + + def _responses_url(self) -> str: + return f"{(self.base_url or CODEX_BASE_URL).rstrip('/')}/responses" + + def _build_request_body(self, prompt: str) -> dict[str, object]: + return { + "model": self.model, + "instructions": "", + "input": [ + { + "type": "message", + "role": "user", + "content": [ + { + "type": "input_text", + "text": prompt, + } + ], + } + ], + "tools": [], + "tool_choice": "auto", + "parallel_tool_calls": False, + "reasoning": None, + "store": False, + "stream": True, + "include": [], + "prompt_cache_key": self.conversation_id, + } + + def _build_headers(self, credentials: _Credentials) -> dict[str, str]: + headers = { + "Authorization": f"Bearer {credentials.access_token}", + "Content-Type": "application/json", + "Accept": "text/event-stream", + "session_id": self.conversation_id, + } + if credentials.account_id: + headers["ChatGPT-Account-ID"] = credentials.account_id + return headers + + def _post_stream( + self, + url: str, + headers: Mapping[str, str], + body: Mapping[str, object], + ) -> Iterable[bytes]: + request = Request( + url, + data=json.dumps(body).encode("utf-8"), + headers=dict(headers), + method="POST", + ) + try: + return urlopen(request, timeout=self.http_timeout) # noqa: S310 - research-only client uses explicit OpenAI/Codex URL. + except HTTPError as exc: + raise _HttpStatusError( + exc.code, f"Codex inference HTTP status {exc.code}" + ) from exc + + def _post_json(self, url: str, body: Mapping[str, object]) -> Any: + request = Request( + url, + data=json.dumps(body).encode("utf-8"), + headers={"Content-Type": "application/json"}, + method="POST", + ) + try: + with urlopen(request, timeout=self.http_timeout) as response: # noqa: S310 - refresh URL is fixed OpenAI OAuth endpoint. + return json.loads(response.read().decode("utf-8")) + except HTTPError as exc: + raise _HttpStatusError( + exc.code, f"Codex refresh HTTP status {exc.code}" + ) from exc + except json.JSONDecodeError as exc: + raise RuntimeError( + "Codex ChatGPT auth refresh returned malformed response" + ) from exc + + def _parse_sse(self, stream: Iterable[bytes]) -> str: + deltas: list[str] = [] + fallback_text: Optional[str] = None + event_lines: list[str] = [] + completed = False + + for raw_line in stream: + line = raw_line.decode("utf-8").rstrip("\r\n") + if not line: + completed = self._process_event_lines( + event_lines, deltas, fallback_text + ) + fallback_text = ( + self._extract_fallback_from_event_lines(event_lines) + or fallback_text + ) + event_lines = [] + if completed: + break + continue + if line.startswith(":"): + continue + event_lines.append(line) + + if event_lines and not completed: + completed = self._process_event_lines(event_lines, deltas, fallback_text) + fallback_text = ( + self._extract_fallback_from_event_lines(event_lines) or fallback_text + ) + + if not completed: + raise _RetryableCodexStreamError( + "Codex SSE stream ended before completion" + ) + + output = "".join(deltas) if deltas else fallback_text + if not output: + raise _RetryableCodexStreamError("Codex completed response had no text") + return output + + def _process_event_lines( + self, + event_lines: list[str], + deltas: list[str], + fallback_text: Optional[str], + ) -> bool: + del fallback_text + if not event_lines: + return False + payload = self._parse_event_payload(event_lines) + event_type = payload.get("type") + if event_type == "response.output_text.delta": + delta = payload.get("delta") + if isinstance(delta, str): + deltas.append(delta) + return False + if event_type == "response.output_item.done": + return False + if event_type in {"response.completed", "response.done"}: + return True + if event_type == "error" or "error" in payload: + raise _RetryableCodexStreamError("Codex SSE stream returned error event") + return False + + def _extract_fallback_from_event_lines( + self, event_lines: list[str] + ) -> Optional[str]: + if not event_lines: + return None + payload = self._parse_event_payload(event_lines) + if payload.get("type") != "response.output_item.done": + return None + item = payload.get("item") + if not isinstance(item, dict) or item.get("type") != "message": + return None + content = item.get("content") + if not isinstance(content, list): + return None + parts: list[str] = [] + for part in content: + if not isinstance(part, dict) or part.get("type") != "output_text": + continue + text = part.get("text") + if isinstance(text, str): + parts.append(text) + return "".join(parts) if parts else None + + def _parse_event_payload(self, event_lines: list[str]) -> dict[str, Any]: + data_lines = [ + line[5:].lstrip(" ") for line in event_lines if line.startswith("data:") + ] + if not data_lines: + return {} + try: + payload = json.loads("\n".join(data_lines)) + except json.JSONDecodeError as exc: + raise _RetryableCodexStreamError( + "Codex SSE payload is invalid JSON" + ) from exc + if not isinstance(payload, dict): + raise _RetryableCodexStreamError("Codex SSE payload is invalid JSON") + return payload + + def _is_retryable_status(self, status_code: int) -> bool: + return status_code in {408, 409, 429} or status_code >= 500 diff --git a/tests/test_dataset_generation.py b/tests/test_dataset_generation.py index 4518709..4d587bd 100644 --- a/tests/test_dataset_generation.py +++ b/tests/test_dataset_generation.py @@ -1,13 +1,21 @@ import json +import logging import tempfile +import threading import unittest from argparse import Namespace +from http.server import BaseHTTPRequestHandler, ThreadingHTTPServer from pathlib import Path - -import httpx -from openai import BadRequestError - -from ai_research.dataset_generation.cli import run_filter, run_generate, run_transform +from unittest.mock import patch + +from ai_research.dataset_generation.cli import ( + LOG_DATE_FORMAT, + LOG_FORMAT, + main, + run_filter, + run_generate, + run_transform, +) from ai_research.dataset_generation.domain.assertion import parse_assertion from ai_research.dataset_generation.domain.example import ( parse_example, @@ -157,42 +165,41 @@ def infer(self, template, data): return self._responses.pop(0) -class _FakeCompletions: - def __init__(self, responses): - self._responses = list(responses) - self.calls = [] - - def create(self, **kwargs): - self.calls.append(kwargs) - response = self._responses.pop(0) - if isinstance(response, Exception): - raise response - return response - - -class _FakeChat: - def __init__(self, completions): - self.completions = completions - - -class _FakeOpenAI: - def __init__(self, completions): - self.chat = _FakeChat(completions) - - -class _FakeMessage: - def __init__(self, content): - self.content = content +class _RecordingOpenAIServer: + def __init__(self): + self.requests = [] + fixture = self + + class Handler(BaseHTTPRequestHandler): + def do_POST(self): + length = int(self.headers.get("Content-Length", "0")) + body = json.loads(self.rfile.read(length).decode("utf-8")) + fixture.requests.append( + {"path": self.path, "headers": dict(self.headers), "json": body} + ) + payload = {"choices": [{"message": {"content": "ok"}}]} + response = json.dumps(payload).encode("utf-8") + self.send_response(200) + self.send_header("Content-Type", "application/json") + self.send_header("Content-Length", str(len(response))) + self.end_headers() + self.wfile.write(response) + def log_message(self, format, *args): + pass -class _FakeChoice: - def __init__(self, content): - self.message = _FakeMessage(content) + self._server = ThreadingHTTPServer(("127.0.0.1", 0), Handler) + self.url = f"http://127.0.0.1:{self._server.server_port}/v1" + self._thread = threading.Thread(target=self._server.serve_forever, daemon=True) + def __enter__(self): + self._thread.start() + return self -class _FakeResponse: - def __init__(self, content): - self.choices = [_FakeChoice(content)] + def __exit__(self, exc_type, exc, traceback): + self._server.shutdown() + self._server.server_close() + self._thread.join(timeout=2) class DatasetGenerationTests(unittest.TestCase): @@ -325,8 +332,6 @@ class Helper: self.assertIn('snippet = "### Example fake heading"', sections[0]) def test_parse_examples_propagates_non_value_error(self): - from unittest.mock import patch - with patch( "ai_research.dataset_generation.domain.llm_response.parse_example", side_effect=RuntimeError("boom"), @@ -451,34 +456,149 @@ def test_run_transform_flattens_grouped_jsonl(self): self.assertEqual(json.loads(rows[0])["assertion"]["assertionText"], "one") self.assertEqual(json.loads(rows[1])["assertion"]["assertionText"], "two") - def test_openai_client_retries_with_max_completion_tokens(self): - client = OpenAITextClient( - api_key="test-key", - model="gpt-5-mini", - temperature=1.0, - max_tokens=123, - ) - bad_request = BadRequestError( - "Unsupported parameter: 'max_tokens' is not supported with this model. Use 'max_completion_tokens' instead.", - response=httpx.Response( - 400, - request=httpx.Request( - "POST", "https://api.openai.com/v1/chat/completions" - ), - ), - body={ - "error": { - "message": "Unsupported parameter: 'max_tokens' is not supported with this model. Use 'max_completion_tokens' instead." - } - }, - ) - completions = _FakeCompletions([bad_request, _FakeResponse("ok")]) - client._client = _FakeOpenAI(completions) + def test_main_configures_timestamped_logging(self): + with tempfile.TemporaryDirectory() as tmp_dir: + input_path = Path(tmp_dir) / "output.jsonl" + output_path = Path(tmp_dir) / "train.jsonl" + input_path.write_text("") + + root_logger = logging.getLogger() + old_handlers = list(root_logger.handlers) + old_level = root_logger.level + try: + root_logger.handlers.clear() + result = main( + ["transform", "-i", str(input_path), "-o", str(output_path)] + ) - result = client.infer("Hello {{.Name}}", {"Name": "world"}) + self.assertEqual(result, 0) + self.assertEqual(root_logger.level, logging.INFO) + self.assertTrue(root_logger.handlers) + formatter = root_logger.handlers[0].formatter + self.assertIsNotNone(formatter) + self.assertEqual(formatter._fmt, LOG_FORMAT) + self.assertEqual(formatter.datefmt, LOG_DATE_FORMAT) + finally: + root_logger.handlers[:] = old_handlers + root_logger.setLevel(old_level) + + def test_run_transform_rejects_misparsed_examples_with_prompt_artifacts(self): + with tempfile.TemporaryDirectory() as tmp_dir: + input_path = Path(tmp_dir) / "output.jsonl" + output_path = Path(tmp_dir) / "train.jsonl" + valid_example = { + "assertion": {"assertionText": "valid", "codeObjectNames": []}, + "codeObjects": [], + "thoughts": "valid", + "explanation": None, + } + input_path.write_text( + json.dumps( + { + "category": "Business Logic", + "examples": [ + valid_example, + { + "assertion": { + "assertionText": "leaked [Assertion] marker", + "codeObjectNames": [], + }, + "codeObjects": [], + "thoughts": "invalid", + "explanation": None, + }, + { + "assertion": { + "assertionText": "invalid", + "codeObjectNames": [], + }, + "codeObjects": [ + {"name": "Service", "code": "leaked [Code] marker"} + ], + "thoughts": "invalid", + "explanation": None, + }, + { + "assertion": { + "assertionText": "invalid", + "codeObjectNames": [], + }, + "codeObjects": [], + "thoughts": ["leaked [Thinking] marker"], + "explanation": None, + }, + { + "assertion": { + "assertionText": "invalid", + "codeObjectNames": [], + }, + "codeObjects": [], + "thoughts": "invalid", + "explanation": {"details": "leaked Example Set marker"}, + }, + { + "assertion": { + "assertionText": "invalid", + "codeObjectNames": [], + }, + "codeObjects": [], + "thoughts": "invalid", + "explanation": "leaked [Explanation] marker", + }, + { + "assertion": { + "assertionText": "invalid", + "codeObjectNames": [], + }, + "codeObjects": [], + "thoughts": "leaked [assertion] marker", + "explanation": None, + }, + { + "assertion": { + "assertionText": "invalid", + "codeObjectNames": [], + }, + "codeObjects": [], + "thoughts": {"nested": ["leaked [explanation] marker"]}, + "explanation": None, + }, + ], + } + ) + + "\n" + ) + args = Namespace(input=str(input_path), output=str(output_path)) + + with self.assertLogs( + "ai_research.dataset_generation.cli", level="WARNING" + ) as logs: + run_transform(args) + + rows = output_path.read_text().splitlines() + self.assertEqual([json.loads(row) for row in rows], [valid_example]) + self.assertEqual(len(logs.output), 7) + self.assertIn("line 1, example 2", logs.output[0]) + self.assertIn("line 1, example 8", logs.output[6]) + + def test_openai_client_sends_prompt_without_token_limit(self): + with _RecordingOpenAIServer() as server: + client = OpenAITextClient( + api_key="test-key", + base_url=server.url, + model="gpt-5-mini", + temperature=1.0, + max_tokens=123, + ) + + result = client.infer("Hello {{.Name}}", {"Name": "world"}) self.assertEqual(result, "ok") - self.assertEqual(completions.calls[0]["max_tokens"], 123) - self.assertNotIn("max_completion_tokens", completions.calls[0]) - self.assertEqual(completions.calls[1]["max_completion_tokens"], 123) - self.assertNotIn("max_tokens", completions.calls[1]) + request_json = server.requests[0]["json"] + self.assertEqual(request_json["model"], "gpt-5-mini") + self.assertEqual( + request_json["messages"], + [{"role": "user", "content": "Hello world"}], + ) + self.assertNotIn("max_tokens", request_json) + self.assertNotIn("max_completion_tokens", request_json) diff --git a/tests/test_openai_codex_client.py b/tests/test_openai_codex_client.py new file mode 100644 index 0000000..46796b3 --- /dev/null +++ b/tests/test_openai_codex_client.py @@ -0,0 +1,647 @@ +import argparse +import json +import os +import tempfile +import threading +import unittest +from collections.abc import Iterable, Iterator, Mapping +from datetime import datetime, timedelta, timezone +from http.server import BaseHTTPRequestHandler, ThreadingHTTPServer +from pathlib import Path +from typing import Any, Optional, cast +from unittest.mock import patch + +from ai_research.dataset_generation import infrastructure +from ai_research.dataset_generation.cli import _build_client +from ai_research.dataset_generation.infrastructure.openai_client import ( + OpenAITextClient, + render_prompt_template as render_openai_prompt_template, +) +from ai_research.dataset_generation.infrastructure.openai_codex_client import ( + OpenAICodexTextClient, + render_prompt_template, +) + + +def _encode_sse(*payloads: Mapping[str, object]) -> bytes: + lines: list[bytes] = [] + for payload in payloads: + lines.append(f"data: {json.dumps(payload)}\n".encode()) + lines.append(b"\n") + return b"".join(lines) + + +def _write_auth( + home: Path, + *, + access_token: str = "access-token", + refresh_token: Optional[str] = "refresh-token", + account_id: Optional[str] = "account-id", + last_refresh: Optional[str] = None, + auth_mode: Optional[str] = "chatgpt", +) -> None: + tokens: dict[str, object] = { + "id_token": "id-token", + "access_token": access_token, + "refresh_token": refresh_token, + } + if account_id is not None: + tokens["account_id"] = account_id + payload: dict[str, object] = { + "OPENAI_API_KEY": None, + "tokens": tokens, + } + if auth_mode is not None: + payload["auth_mode"] = auth_mode + if last_refresh is not None: + payload["last_refresh"] = last_refresh + home.mkdir(exist_ok=True) + (home / "auth.json").write_text(json.dumps(payload)) + + +def _header(headers: Mapping[str, str], name: str) -> str: + for key, value in headers.items(): + if key.lower() == name.lower(): + return value + raise KeyError(name) + + +class _RecordingCodexServer: + def __init__(self) -> None: + self.requests: list[dict[str, Any]] = [] + self.refresh_requests: list[dict[str, Any]] = [] + self.responses: list[tuple[int, bytes, str]] = [] + self.refresh_responses: list[tuple[int, bytes, str]] = [] + fixture = self + + class Handler(BaseHTTPRequestHandler): + def do_POST(self) -> None: + length = int(self.headers.get("Content-Length", "0")) + raw_body = self.rfile.read(length).decode("utf-8") + body = json.loads(raw_body) if raw_body else {} + record = { + "path": self.path, + "headers": dict(self.headers), + "json": body, + } + + if self.path == "/refresh": + fixture.refresh_requests.append(record) + status, response, content_type = fixture.refresh_responses.pop(0) + else: + fixture.requests.append(record) + status, response, content_type = fixture.responses.pop(0) + + self.send_response(status) + self.send_header("Content-Type", content_type) + self.send_header("Content-Length", str(len(response))) + self.end_headers() + self.wfile.write(response) + + def log_message(self, format: str, *args: object) -> None: + pass + + self._server = ThreadingHTTPServer(("127.0.0.1", 0), Handler) + self.url = f"http://127.0.0.1:{self._server.server_port}" + self.refresh_url = f"{self.url}/refresh" + self._thread = threading.Thread(target=self._server.serve_forever, daemon=True) + + def __enter__(self) -> "_RecordingCodexServer": + self._thread.start() + return self + + def __exit__(self, exc_type: object, exc: object, traceback: object) -> None: + self._server.shutdown() + self._server.server_close() + self._thread.join(timeout=2) + + def add_sse(self, *payloads: Mapping[str, object], status: int = 200) -> None: + self.responses.append((status, _encode_sse(*payloads), "text/event-stream")) + + def add_incomplete_sse(self, *payloads: Mapping[str, object]) -> None: + self.responses.append((200, _encode_sse(*payloads), "text/event-stream")) + + def add_status(self, status: int, body: bytes = b"error") -> None: + self.responses.append((status, body, "text/plain")) + + def add_refresh_json( + self, payload: Mapping[str, object], status: int = 200 + ) -> None: + self.refresh_responses.append( + (status, json.dumps(payload).encode("utf-8"), "application/json") + ) + + +class _CloseAwareStream: + def __init__(self, lines: Iterable[bytes]) -> None: + self._lines = list(lines) + self.closed = False + + def __iter__(self) -> Iterator[bytes]: + return iter(self._lines) + + def close(self) -> None: + self.closed = True + + +class _CloseAwareCodexClient(OpenAICodexTextClient): + def __init__(self, stream: _CloseAwareStream) -> None: + super().__init__(api_key="", model="gpt-5-codex", temperature=1.0, max_tokens=1) + self.stream = stream + + def _post_stream( + self, + url: str, + headers: Mapping[str, str], + body: Mapping[str, object], + ) -> Iterable[bytes]: + del url, headers, body + return self.stream + + +class OpenAICodexTextClientTests(unittest.TestCase): + def setUp(self) -> None: + self._old_codex_home = os.environ.get("CODEX_HOME") + + def tearDown(self) -> None: + if self._old_codex_home is None: + os.environ.pop("CODEX_HOME", None) + else: + os.environ["CODEX_HOME"] = self._old_codex_home + + def test_template_rendering_matches_openai_client(self) -> None: + template = "Hello {{.Name}}, count {{.Count}}" + data = {"Name": "world", "Count": 3} + + self.assertEqual( + render_prompt_template(template, data), + render_openai_prompt_template(template, data), + ) + + def test_constructor_stores_compatibility_fields(self) -> None: + client = OpenAICodexTextClient( + api_key="", + base_url="https://example.test/root", + model="gpt-5-codex", + temperature=0.5, + max_tokens=456, + max_retries=2, + retry_delay_seconds=0.25, + http_timeout=3.5, + ) + + self.assertEqual(client.api_key, "") + self.assertEqual(client.base_url, "https://example.test/root") + self.assertEqual(client.model, "gpt-5-codex") + self.assertEqual(client.temperature, 0.5) + self.assertEqual(client.max_tokens, 456) + self.assertEqual(client.max_retries, 2) + self.assertEqual(client.retry_delay_seconds, 0.25) + self.assertEqual(client.http_timeout, 3.5) + + def test_constructor_rejects_non_empty_api_key(self) -> None: + with self.assertRaisesRegex(ValueError, "use OpenAITextClient"): + OpenAICodexTextClient( + api_key="key", model="gpt-5-codex", temperature=1.0, max_tokens=1 + ) + + def test_cli_build_client_uses_openai_for_non_empty_api_key(self) -> None: + args = argparse.Namespace( + api_key="key", + base_url=None, + model="gpt-4o-mini", + temperature=1.0, + max_tokens=1, + ) + + self.assertIsInstance(_build_client(args), OpenAITextClient) + + def test_cli_build_client_uses_codex_for_empty_api_key(self) -> None: + args = argparse.Namespace( + api_key="", + base_url=None, + model="gpt-5-codex", + temperature=1.0, + max_tokens=1, + ) + + self.assertIsInstance(_build_client(args), OpenAICodexTextClient) + + def test_infer_loads_chatgpt_file_auth_and_uses_http_fixture(self) -> None: + with ( + tempfile.TemporaryDirectory() as tmp_dir, + _RecordingCodexServer() as server, + ): + os.environ["CODEX_HOME"] = tmp_dir + _write_auth(Path(tmp_dir)) + server.add_sse( + {"type": "response.output_text.delta", "delta": "ok"}, + {"type": "response.completed"}, + ) + client = OpenAICodexTextClient( + api_key="", + base_url=server.url, + model="gpt-5-codex", + temperature=1.0, + max_tokens=123, + ) + + self.assertEqual(client.infer("Hi {{.Name}}", {"Name": "there"}), "ok") + + request = server.requests[0] + self.assertEqual(request["path"], "/responses") + headers = request["headers"] + self.assertEqual(_header(headers, "Authorization"), "Bearer access-token") + self.assertEqual(_header(headers, "Content-Type"), "application/json") + self.assertEqual(_header(headers, "Accept"), "text/event-stream") + self.assertEqual(_header(headers, "session_id"), client.conversation_id) + self.assertEqual(_header(headers, "ChatGPT-Account-ID"), "account-id") + input_items = cast(list[dict[str, Any]], request["json"]["input"]) + content_items = cast(list[dict[str, Any]], input_items[0]["content"]) + self.assertEqual(content_items[0]["text"], "Hi there") + + def test_infer_accepts_missing_auth_mode_when_tokens_are_present(self) -> None: + with ( + tempfile.TemporaryDirectory() as tmp_dir, + _RecordingCodexServer() as server, + ): + os.environ["CODEX_HOME"] = tmp_dir + _write_auth(Path(tmp_dir), auth_mode=None) + server.add_sse( + {"type": "response.output_text.delta", "delta": "ok"}, + {"type": "response.completed"}, + ) + client = OpenAICodexTextClient( + api_key="", + base_url=server.url, + model="gpt-5-codex", + temperature=1.0, + max_tokens=1, + ) + + self.assertEqual(client.infer("Hi", {}), "ok") + + def test_base_url_overrides_codex_endpoint_root(self) -> None: + with ( + tempfile.TemporaryDirectory() as tmp_dir, + _RecordingCodexServer() as server, + ): + os.environ["CODEX_HOME"] = tmp_dir + _write_auth(Path(tmp_dir), account_id=None) + server.add_sse( + {"type": "response.output_text.delta", "delta": "ok"}, + {"type": "response.done"}, + ) + client = OpenAICodexTextClient( + api_key="", + base_url=f"{server.url}/codex/", + model="gpt-5-codex", + temperature=1.0, + max_tokens=1, + ) + + client.infer("Hi", {}) + + self.assertEqual(server.requests[0]["path"], "/codex/responses") + self.assertNotIn("ChatGPT-Account-ID", server.requests[0]["headers"]) + + def test_request_body_contains_no_tool_responses_shape(self) -> None: + with ( + tempfile.TemporaryDirectory() as tmp_dir, + _RecordingCodexServer() as server, + ): + os.environ["CODEX_HOME"] = tmp_dir + _write_auth(Path(tmp_dir)) + server.add_sse( + {"type": "response.output_text.delta", "delta": "ok"}, + {"type": "response.completed"}, + ) + client = OpenAICodexTextClient( + api_key="", + base_url=server.url, + model="gpt-5-codex", + temperature=1.0, + max_tokens=123, + ) + + client.infer("Prompt", {}) + + body = server.requests[0]["json"] + self.assertEqual(body["model"], "gpt-5-codex") + self.assertEqual(body["instructions"], "") + self.assertEqual(body["tools"], []) + self.assertEqual(body["tool_choice"], "auto") + self.assertFalse(body["parallel_tool_calls"]) + self.assertIsNone(body["reasoning"]) + self.assertFalse(body["store"]) + self.assertTrue(body["stream"]) + self.assertEqual(body["include"], []) + self.assertEqual(body["prompt_cache_key"], client.conversation_id) + self.assertNotIn("temperature", body) + self.assertNotIn("max_tokens", body) + + def test_infer_joins_output_deltas_and_accepts_done(self) -> None: + with ( + tempfile.TemporaryDirectory() as tmp_dir, + _RecordingCodexServer() as server, + ): + os.environ["CODEX_HOME"] = tmp_dir + _write_auth(Path(tmp_dir)) + server.add_sse( + {"type": "response.output_text.delta", "delta": "hello"}, + {"type": "response.output_text.delta", "delta": " world"}, + {"type": "response.done"}, + ) + client = OpenAICodexTextClient( + api_key="", + base_url=server.url, + model="gpt-5-codex", + temperature=1.0, + max_tokens=1, + ) + + self.assertEqual(client.infer("Prompt", {}), "hello world") + + def test_infer_falls_back_to_output_item_done_text(self) -> None: + with ( + tempfile.TemporaryDirectory() as tmp_dir, + _RecordingCodexServer() as server, + ): + os.environ["CODEX_HOME"] = tmp_dir + _write_auth(Path(tmp_dir)) + server.add_sse( + { + "type": "response.output_item.done", + "item": { + "type": "message", + "content": [{"type": "output_text", "text": "full text"}], + }, + }, + {"type": "response.completed"}, + ) + client = OpenAICodexTextClient( + api_key="", + base_url=server.url, + model="gpt-5-codex", + temperature=1.0, + max_tokens=1, + ) + + self.assertEqual(client.infer("Prompt", {}), "full text") + + def test_transient_response_retries_with_exponential_backoff(self) -> None: + with ( + tempfile.TemporaryDirectory() as tmp_dir, + _RecordingCodexServer() as server, + ): + os.environ["CODEX_HOME"] = tmp_dir + _write_auth(Path(tmp_dir)) + server.add_status(500) + server.add_status(429) + server.add_sse( + {"type": "response.output_text.delta", "delta": "ok"}, + {"type": "response.completed"}, + ) + client = OpenAICodexTextClient( + api_key="", + base_url=server.url, + model="gpt-5-codex", + temperature=1.0, + max_tokens=1, + ) + sleeps: list[float] = [] + client._sleep = sleeps.append + + self.assertEqual(client.infer("Prompt", {}), "ok") + self.assertEqual(len(server.requests), 3) + self.assertEqual(sleeps, [1.0, 2.0]) + + def test_incomplete_sse_stream_retries_with_exponential_backoff(self) -> None: + with ( + tempfile.TemporaryDirectory() as tmp_dir, + _RecordingCodexServer() as server, + ): + os.environ["CODEX_HOME"] = tmp_dir + _write_auth(Path(tmp_dir)) + server.add_incomplete_sse( + {"type": "response.output_text.delta", "delta": "partial"} + ) + server.add_sse( + {"type": "response.output_text.delta", "delta": "ok"}, + {"type": "response.completed"}, + ) + client = OpenAICodexTextClient( + api_key="", + base_url=server.url, + model="gpt-5-codex", + temperature=1.0, + max_tokens=1, + ) + sleeps: list[float] = [] + client._sleep = sleeps.append + + self.assertEqual(client.infer("Prompt", {}), "ok") + self.assertEqual(len(server.requests), 2) + self.assertEqual(sleeps, [1.0]) + + def test_incomplete_sse_stream_aborts_after_exhausting_retries(self) -> None: + with ( + tempfile.TemporaryDirectory() as tmp_dir, + _RecordingCodexServer() as server, + ): + os.environ["CODEX_HOME"] = tmp_dir + _write_auth(Path(tmp_dir)) + for delta in ("one", "two", "three"): + server.add_incomplete_sse( + {"type": "response.output_text.delta", "delta": delta} + ) + client = OpenAICodexTextClient( + api_key="", + base_url=server.url, + model="gpt-5-codex", + temperature=1.0, + max_tokens=1, + max_retries=2, + ) + sleeps: list[float] = [] + client._sleep = sleeps.append + + with self.assertRaisesRegex( + RuntimeError, "Codex inference exhausted retries" + ): + client.infer("Prompt", {}) + self.assertEqual(len(server.requests), 3) + self.assertEqual(sleeps, [1.0, 2.0]) + + def test_non_transient_400_does_not_retry(self) -> None: + with ( + tempfile.TemporaryDirectory() as tmp_dir, + _RecordingCodexServer() as server, + ): + os.environ["CODEX_HOME"] = tmp_dir + _write_auth(Path(tmp_dir)) + server.add_status(400) + client = OpenAICodexTextClient( + api_key="", + base_url=server.url, + model="gpt-5-codex", + temperature=1.0, + max_tokens=1, + ) + sleeps: list[float] = [] + client._sleep = sleeps.append + + with self.assertRaisesRegex(RuntimeError, "non-retryable HTTP status 400"): + client.infer("Prompt", {}) + self.assertEqual(len(server.requests), 1) + self.assertEqual(sleeps, []) + + def test_stale_chatgpt_auth_refreshes_before_inference_and_preserves_account_id( + self, + ) -> None: + with ( + tempfile.TemporaryDirectory() as tmp_dir, + _RecordingCodexServer() as server, + ): + os.environ["CODEX_HOME"] = tmp_dir + old_refresh = (datetime.now(timezone.utc) - timedelta(days=9)).isoformat() + _write_auth(Path(tmp_dir), last_refresh=old_refresh) + server.add_refresh_json( + {"access_token": "new-access", "refresh_token": "new-refresh"} + ) + server.add_sse( + {"type": "response.output_text.delta", "delta": "ok"}, + {"type": "response.completed"}, + ) + client = OpenAICodexTextClient( + api_key="", + base_url=server.url, + model="gpt-5-codex", + temperature=1.0, + max_tokens=1, + ) + + with patch.object( + infrastructure.openai_codex_client, "REFRESH_URL", server.refresh_url + ): + client.infer("Prompt", {}) + + self.assertEqual(len(server.refresh_requests), 1) + self.assertEqual( + _header(server.requests[0]["headers"], "Authorization"), + "Bearer new-access", + ) + auth_payload = json.loads((Path(tmp_dir) / "auth.json").read_text()) + self.assertEqual(auth_payload["tokens"]["account_id"], "account-id") + self.assertEqual(auth_payload["tokens"]["refresh_token"], "new-refresh") + + def test_first_401_refreshes_once_and_retries_once(self) -> None: + with ( + tempfile.TemporaryDirectory() as tmp_dir, + _RecordingCodexServer() as server, + ): + os.environ["CODEX_HOME"] = tmp_dir + _write_auth(Path(tmp_dir)) + server.add_status(401) + server.add_refresh_json( + {"access_token": "new-access", "refresh_token": "new-refresh"} + ) + server.add_sse( + {"type": "response.output_text.delta", "delta": "ok"}, + {"type": "response.completed"}, + ) + client = OpenAICodexTextClient( + api_key="", + base_url=server.url, + model="gpt-5-codex", + temperature=1.0, + max_tokens=1, + ) + + with patch.object( + infrastructure.openai_codex_client, "REFRESH_URL", server.refresh_url + ): + self.assertEqual(client.infer("Prompt", {}), "ok") + + self.assertEqual(len(server.refresh_requests), 1) + self.assertEqual(len(server.requests), 2) + self.assertEqual( + _header(server.requests[1]["headers"], "Authorization"), + "Bearer new-access", + ) + + def test_errors_are_sanitized(self) -> None: + with ( + tempfile.TemporaryDirectory() as tmp_dir, + _RecordingCodexServer() as server, + ): + os.environ["CODEX_HOME"] = tmp_dir + _write_auth( + Path(tmp_dir), + access_token="secret-access", + refresh_token="secret-refresh", + ) + server.add_status(400, b"secret-access secret-refresh") + client = OpenAICodexTextClient( + api_key="", + base_url=server.url, + model="gpt-5-codex", + temperature=1.0, + max_tokens=1, + ) + + with self.assertRaises(RuntimeError) as exc: + client.infer("Prompt", {}) + + message = str(exc.exception) + self.assertNotIn("secret-access", message) + self.assertNotIn("secret-refresh", message) + + def test_refresh_missing_access_token_does_not_change_auth_file(self) -> None: + with ( + tempfile.TemporaryDirectory() as tmp_dir, + _RecordingCodexServer() as server, + ): + os.environ["CODEX_HOME"] = tmp_dir + old_refresh = (datetime.now(timezone.utc) - timedelta(days=9)).isoformat() + _write_auth(Path(tmp_dir), last_refresh=old_refresh) + auth_path = Path(tmp_dir) / "auth.json" + original_auth = auth_path.read_text() + server.add_refresh_json({"refresh_token": "new-refresh"}) + client = OpenAICodexTextClient( + api_key="", + base_url=server.url, + model="gpt-5-codex", + temperature=1.0, + max_tokens=1, + ) + + with patch.object( + infrastructure.openai_codex_client, "REFRESH_URL", server.refresh_url + ): + with self.assertRaisesRegex(RuntimeError, "returned no access token"): + client.infer("Prompt", {}) + + self.assertEqual(auth_path.read_text(), original_auth) + self.assertEqual(server.requests, []) + + def test_stream_is_closed_after_infer(self) -> None: + with tempfile.TemporaryDirectory() as tmp_dir: + os.environ["CODEX_HOME"] = tmp_dir + _write_auth(Path(tmp_dir)) + stream = _CloseAwareStream( + [ + *_encode_sse( + {"type": "response.output_text.delta", "delta": "ok"}, + {"type": "response.completed"}, + ).splitlines(keepends=True) + ] + ) + client = _CloseAwareCodexClient(stream) + + self.assertEqual(client.infer("Prompt", {}), "ok") + self.assertTrue(stream.closed) + + +if __name__ == "__main__": + unittest.main()