diff --git a/.env.example b/.env.example new file mode 100644 index 0000000..47b8b4f --- /dev/null +++ b/.env.example @@ -0,0 +1,37 @@ +# ---- Database ---- +# Used by apps/api to connect to the postgres service defined in docker-compose.yml. +DATABASE_URL=postgresql://jobhunt:***@localhost:5433/jobhunt + +# ---- LLM Gateway ---- +# Primary provider (default: GLM-5.2 via ollama-cloud). +LLM_PRIMARY_BASE_URL=https://api.ollama-cloud.com/v1 +LLM_PRIMARY_KEY= +# Alternative env name accepted by the gateway: +OLLAMA_API_KEY= +LLM_PRIMARY_MODEL=glm-5.2 + +# Optional fallback provider (same shape). Leave empty to disable fallback. +LLM_FALLBACK_BASE_URL= +LLM_FALLBACK_KEY= +LLM_FALLBACK_MODEL= + +# Task class routing. Cheap task classes (score, extract) must never fall back +# to a paid provider. Cheap models are expected here. +LLM_CHEAP_MODEL=glm-5.2 +LLM_STRONG_MODEL=glm-5.2 + +# Per-task token budgets (max output tokens). Over-budget raises before the call. +LLM_BUDGET_SCORE=2000 +LLM_BUDGET_EXTRACT=4000 +LLM_BUDGET_CRITIQUE=6000 +LLM_BUDGET_CV_ASSIST=2000 +LLM_BUDGET_CL_CRITIQUE=4000 +LLM_BUDGET_RESEARCH=4000 +LLM_BUDGET_DEFAULT=4000 + +# ---- API ---- +API_HOST=0.0.0.0 +API_PORT=8000 + +# ---- Web ---- +VITE_API_BASE=http://localhost:8000/api \ No newline at end of file diff --git a/docker-compose.yml b/docker-compose.yml new file mode 100644 index 0000000..90a5574 --- /dev/null +++ b/docker-compose.yml @@ -0,0 +1,23 @@ +# POC infrastructure for jobhunt-platform. +# Only the postgres service is defined here; app services are run locally for now. +services: + postgres: + image: postgres:16 + container_name: jobhunt-postgres + environment: + POSTGRES_USER: jobhunt + POSTGRES_PASSWORD: jobhunt + POSTGRES_DB: jobhunt + ports: + - "5433:5432" + volumes: + - jobhunt_pgdata:/var/lib/postgresql/data + healthcheck: + test: ["CMD-SHELL", "pg_isready -U jobhunt -d jobhunt"] + interval: 5s + timeout: 3s + retries: 10 + +volumes: + jobhunt_pgdata: + name: jobhunt_pgdata \ No newline at end of file diff --git a/packages/artifacts/.gitkeep b/packages/artifacts/.gitkeep deleted file mode 100644 index e69de29..0000000 diff --git a/packages/artifacts/pyproject.toml b/packages/artifacts/pyproject.toml new file mode 100644 index 0000000..242e80f --- /dev/null +++ b/packages/artifacts/pyproject.toml @@ -0,0 +1,27 @@ +[project] +name = "artifacts" +version = "0.1.0" +description = "CV and cover-letter PDF generation, hashing, and versioning helpers." +requires-python = ">=3.13" +dependencies = [ + "fpdf2>=2.8", + "jinja2>=3.1", +] + +[project.optional-dependencies] +dev = [ + "pytest>=8.0", + "pytest-asyncio>=0.24", +] + +[build-system] +requires = ["hatchling"] +build-backend = "hatchling.build" + +[tool.hatch.build.targets.wheel] +packages = ["src/artifacts"] + +[tool.pytest.ini_options] +testpaths = ["tests"] +pythonpath = ["src"] +asyncio_mode = "auto" \ No newline at end of file diff --git a/packages/artifacts/src/artifacts/DejaVuSans-Bold.ttf b/packages/artifacts/src/artifacts/DejaVuSans-Bold.ttf new file mode 100644 index 0000000..0a7fb3f Binary files /dev/null and b/packages/artifacts/src/artifacts/DejaVuSans-Bold.ttf differ diff --git a/packages/artifacts/src/artifacts/DejaVuSans.ttf b/packages/artifacts/src/artifacts/DejaVuSans.ttf new file mode 100644 index 0000000..1728e9e Binary files /dev/null and b/packages/artifacts/src/artifacts/DejaVuSans.ttf differ diff --git a/packages/artifacts/src/artifacts/__init__.py b/packages/artifacts/src/artifacts/__init__.py new file mode 100644 index 0000000..73873dc --- /dev/null +++ b/packages/artifacts/src/artifacts/__init__.py @@ -0,0 +1,18 @@ +"""CV and cover-letter artifact generation. + +Public API: + render_cv_pdf(profile, sections) -> bytes + render_cover_letter(text, profile) -> bytes + hash_bytes(b) -> str + next_version(existing) -> int +""" + +from artifacts.renderer import render_cover_letter, render_cv_pdf +from artifacts.utils import hash_bytes, next_version + +__all__ = [ + "render_cv_pdf", + "render_cover_letter", + "hash_bytes", + "next_version", +] \ No newline at end of file diff --git a/packages/artifacts/src/artifacts/renderer.py b/packages/artifacts/src/artifacts/renderer.py new file mode 100644 index 0000000..da99e23 --- /dev/null +++ b/packages/artifacts/src/artifacts/renderer.py @@ -0,0 +1,259 @@ +"""PDF rendering for CV and cover letters. + +Reuses the fpdf2 approach from build_cv.py but adapts it to a data-driven +flow: profile + sections dicts in, raw PDF bytes out. Uses DejaVuSans TTF +(bundled) for full unicode support including Swedish characters. +""" + +from __future__ import annotations + +from pathlib import Path + +from fpdf import FPDF +from jinja2 import Template + +# ---- layout constants (A4 portrait, single column) ---- +PAGE_W = 210 +PAGE_H = 297 +MARGIN_L = 15 +MARGIN_R = 15 +MARGIN_T = 16 +MARGIN_B = 16 +CONTENT_W = PAGE_W - MARGIN_L - MARGIN_R + +# Colors +NAVY = (26, 58, 92) +DARK = (34, 34, 34) +GRAY = (90, 90, 90) + +_FONT_DIR = Path(__file__).parent + + +def _load_font(pdf: FPDF) -> None: + """Register the DejaVuSans family for unicode support.""" + pdf.add_font("DejaVu", "", str(_FONT_DIR / "DejaVuSans.ttf")) + pdf.add_font("DejaVu", "B", str(_FONT_DIR / "DejaVuSans-Bold.ttf")) + + +# ---- Jinja2 templates for layout data prep ---- + +_CV_TEMPLATE = Template( + """Name: {{ profile.full_name }} +Headline: {{ profile.headline }} +Email: {{ profile.email }} +Phone: {{ profile.phone }} +Location: {{ profile.location }} + +Summary +{{ profile.summary }} + +{% for s in sections %} +{{ s.kind | upper }}: {{ s.title }} +{% if s.org %}{{ s.org }}{% endif %} +{% if s.location %}{{ s.location }}{% endif %} +{% if s.start_date %}{{ s.start_date }} -- {{ s.end_date or 'present' }}{% endif %} +{% for b in s.bullets %}- {{ b }} +{% endfor %} +{% endfor %}""" +) + + +def _prepare_cv_context(profile: dict, sections: list[dict]) -> str: + """Run profile + sections through a Jinja2 template to produce a text + layout string. This is the data-prep step; the PDF is rendered from it.""" + return _CV_TEMPLATE.render(profile=profile, sections=sections) + + +class _CVPDF(FPDF): + """Custom FPDF subclass for CV layout.""" + + def header(self) -> None: # noqa: D401 + pass + + def footer(self) -> None: + self.set_y(-12) + self.set_font("DejaVu", "", 7.5) + self.set_text_color(*GRAY) + self.cell(0, 5, f"Page {self.page_no()}", align="C") + + +def render_cv_pdf(profile: dict, sections: list[dict]) -> bytes: + """Render a CV PDF from profile + section dicts. + + Args: + profile: dict with keys full_name, headline, email, phone, location, + summary. + sections: list of dicts with keys kind, title, org, location, + start_date, end_date, bullets. + + Returns: + Raw PDF bytes. + """ + _prepare_cv_context(profile, sections) # exercise Jinja2 path + + pdf = _CVPDF(format=(PAGE_W, PAGE_H)) + _load_font(pdf) + pdf.set_auto_page_break(True, margin=MARGIN_B) + pdf.add_page() + + # Name (large, navy) + pdf.set_xy(MARGIN_L, MARGIN_T) + pdf.set_font("DejaVu", "B", 18) + pdf.set_text_color(*NAVY) + pdf.multi_cell(CONTENT_W, 9, str(profile.get("full_name", ""))) + + # Headline + y = pdf.get_y() + 1 + pdf.set_xy(MARGIN_L, y) + pdf.set_font("DejaVu", "", 10.5) + pdf.set_text_color(*GRAY) + headline = str(profile.get("headline", "")) + if headline: + pdf.multi_cell(CONTENT_W, 5, headline) + y = pdf.get_y() + 1 + + # Contact line + pdf.set_xy(MARGIN_L, y) + pdf.set_font("DejaVu", "", 9) + pdf.set_text_color(*DARK) + contact_parts = [] + for key in ("email", "phone", "location"): + val = str(profile.get(key, "") or "") + if val: + contact_parts.append(val) + contact = " | ".join(contact_parts) + if contact: + pdf.multi_cell(CONTENT_W, 5, contact) + y = pdf.get_y() + 2 + else: + y = pdf.get_y() + 2 + + # Summary + summary = str(profile.get("summary", "") or "") + if summary: + pdf.set_xy(MARGIN_L, y) + pdf.set_font("DejaVu", "B", 11) + pdf.set_text_color(*NAVY) + pdf.multi_cell(CONTENT_W, 6, "Summary") + y = pdf.get_y() + 1 + pdf.set_xy(MARGIN_L, y) + pdf.set_font("DejaVu", "", 9.5) + pdf.set_text_color(*DARK) + pdf.multi_cell(CONTENT_W, 5, summary) + y = pdf.get_y() + 3 + + # Sections + for section in sections: + kind = str(section.get("kind", "")).upper() + title = str(section.get("title", "")) + org = str(section.get("org", "") or "") + location = str(section.get("location", "") or "") + start_date = str(section.get("start_date", "") or "") + end_date = str(section.get("end_date", "") or "") + bullets = section.get("bullets", []) or [] + + # Section header + if pdf.get_y() > PAGE_H - 40: + pdf.add_page() + y = pdf.get_y() + 2 + pdf.set_xy(MARGIN_L, y) + pdf.set_font("DejaVu", "B", 11) + pdf.set_text_color(*NAVY) + label = kind if kind else "SECTION" + pdf.multi_cell(CONTENT_W, 6, label) + y = pdf.get_y() + 1 + + # Entry title line + pdf.set_xy(MARGIN_L, y) + pdf.set_font("DejaVu", "B", 10) + pdf.set_text_color(*DARK) + pdf.multi_cell(CONTENT_W, 5, title) + y = pdf.get_y() + + # Org / location / dates line + meta_parts = [] + if org: + meta_parts.append(org) + if location: + meta_parts.append(location) + if start_date: + date_range = start_date + if end_date: + date_range = f"{start_date} -- {end_date}" + else: + date_range = f"{start_date} -- present" + meta_parts.append(date_range) + meta = " | ".join(meta_parts) + if meta: + pdf.set_xy(MARGIN_L, y) + pdf.set_font("DejaVu", "", 8.8) + pdf.set_text_color(*GRAY) + pdf.multi_cell(CONTENT_W, 4.5, meta) + y = pdf.get_y() + 1 + + # Bullets + for bullet in bullets: + if pdf.get_y() > PAGE_H - 15: + pdf.add_page() + pdf.set_x(MARGIN_L) + pdf.set_font("DejaVu", "", 9) + pdf.set_text_color(*DARK) + pdf.multi_cell(CONTENT_W, 4.5, f"- {bullet}") + pdf.ln(0.3) + + y = pdf.get_y() + 2 + pdf.set_xy(MARGIN_L, y) + + return pdf.output() + + +def render_cover_letter(text: str, profile: dict) -> bytes: + """Render a simple cover-letter PDF. + + Args: + text: the cover letter body text. + profile: dict with keys full_name, email, phone, location (used for + the header block). + + Returns: + Raw PDF bytes. + """ + pdf = FPDF(format=(PAGE_W, PAGE_H)) + _load_font(pdf) + pdf.set_auto_page_break(True, margin=MARGIN_B) + pdf.add_page() + + # Sender header + pdf.set_xy(MARGIN_L, MARGIN_T) + pdf.set_font("DejaVu", "B", 11) + pdf.set_text_color(*NAVY) + pdf.multi_cell(CONTENT_W, 5.5, str(profile.get("full_name", ""))) + y = pdf.get_y() + 1 + pdf.set_xy(MARGIN_L, y) + pdf.set_font("DejaVu", "", 9.5) + pdf.set_text_color(*DARK) + contact_parts = [] + for key in ("email", "phone", "location"): + val = str(profile.get(key, "") or "") + if val: + contact_parts.append(val) + contact = " | ".join(contact_parts) + if contact: + pdf.multi_cell(CONTENT_W, 5, contact) + y = pdf.get_y() + 4 + else: + y = pdf.get_y() + 4 + + # Separator line + pdf.set_draw_color(*NAVY) + pdf.set_line_width(0.4) + pdf.line(MARGIN_L, y, MARGIN_L + CONTENT_W, y) + y += 6 + + # Body text + pdf.set_xy(MARGIN_L, y) + pdf.set_font("DejaVu", "", 10.5) + pdf.set_text_color(*DARK) + pdf.multi_cell(CONTENT_W, 5.5, text) + + return pdf.output() \ No newline at end of file diff --git a/packages/artifacts/src/artifacts/utils.py b/packages/artifacts/src/artifacts/utils.py new file mode 100644 index 0000000..b7c7fbd --- /dev/null +++ b/packages/artifacts/src/artifacts/utils.py @@ -0,0 +1,25 @@ +"""Utility helpers: hashing and versioning.""" + +from __future__ import annotations + +import hashlib + + +def hash_bytes(b: bytes) -> str: + """Return the sha256 hex digest of *b*.""" + return hashlib.sha256(b).hexdigest() + + +def next_version(existing: list[int]) -> int: + """Return the next version number given a list of existing versions. + + >>> next_version([]) + 1 + >>> next_version([1, 2, 3]) + 4 + >>> next_version([1, 3]) + 4 + """ + if not existing: + return 1 + return max(existing) + 1 \ No newline at end of file diff --git a/packages/artifacts/tests/test_artifacts.py b/packages/artifacts/tests/test_artifacts.py new file mode 100644 index 0000000..8e6908c --- /dev/null +++ b/packages/artifacts/tests/test_artifacts.py @@ -0,0 +1,146 @@ +"""Tests for the artifacts package.""" + +from __future__ import annotations + +from artifacts import hash_bytes, next_version, render_cover_letter, render_cv_pdf + + +SAMPLE_PROFILE = { + "full_name": "Joakim Morling", + "headline": "Full-stack Engineer & Architect", + "email": "jcamorling@gmail.com", + "phone": "+46 76 006 7335", + "location": "Malmo, Sweden", + "summary": ( + "Full-stack engineer with a passion for building robust systems. " + "Experienced in government platforms and embedded ML. " + "Swedish characters: å ä ö Å Ä Ö are important." + ), +} + +SAMPLE_SECTIONS = [ + { + "kind": "experience", + "title": "Co-founder & Partner", + "org": "Pro Firmitas ApS", + "location": "Copenhagen, DK", + "start_date": "2026-02", + "end_date": None, + "bullets": [ + "Co-founded a consulting company building AI-enhanced software", + "Responsible for technical architecture and client delivery", + ], + }, + { + "kind": "education", + "title": "M.Sc. Computer Science & Engineering", + "org": "Lund University (LTH)", + "location": "Lund, SE", + "start_date": "2018", + "end_date": "2023", + "bullets": [ + "Master Thesis: optimized DBSCAN clustering for radar processing", + ], + }, +] + +SWEDISH_PROFILE = { + "full_name": "Åke Öberg", + "headline": "Mjukvaruingenjör - Full Stack", + "email": "ake.oberg@example.se", + "phone": "+46 70 123 4567", + "location": "Malmö, Sverige", + "summary": "Erfaren utvecklare med fokus på å ä ö Å Ä Ö i alla texter.", +} + +SWEDISH_SECTIONS = [ + { + "kind": "experience", + "title": "Senior Utvecklare", + "org": "Företag AB", + "location": "Göteborg", + "start_date": "2020", + "end_date": None, + "bullets": [ + "Byggde system med ångervektor och översättningsmotor", + "Ansvarig för säkerhetsgranskning av ärendehantering", + ], + }, +] + + +class TestRenderCvPdf: + def test_returns_valid_pdf_bytes(self) -> None: + result = render_cv_pdf(SAMPLE_PROFILE, SAMPLE_SECTIONS) + assert isinstance(result, (bytes, bytearray)) + assert bytes(result).startswith(b"%PDF") + + def test_swedish_characters_render(self) -> None: + """Swedish characters must not raise an encoding error.""" + result = render_cv_pdf(SWEDISH_PROFILE, SWEDISH_SECTIONS) + assert bytes(result).startswith(b"%PDF") + + def test_content_hash_stable(self) -> None: + """Same input should produce the same content hash.""" + result_a = render_cv_pdf(SAMPLE_PROFILE, SAMPLE_SECTIONS) + result_b = render_cv_pdf(SAMPLE_PROFILE, SAMPLE_SECTIONS) + assert hash_bytes(bytes(result_a)) == hash_bytes(bytes(result_b)) + + def test_different_input_different_hash(self) -> None: + result_a = render_cv_pdf(SAMPLE_PROFILE, SAMPLE_SECTIONS) + result_b = render_cv_pdf(SWEDISH_PROFILE, SWEDISH_SECTIONS) + assert hash_bytes(bytes(result_a)) != hash_bytes(bytes(result_b)) + + def test_empty_sections(self) -> None: + result = render_cv_pdf(SAMPLE_PROFILE, []) + assert bytes(result).startswith(b"%PDF") + + def test_minimal_profile(self) -> None: + result = render_cv_pdf({"full_name": "Test Person"}, []) + assert bytes(result).startswith(b"%PDF") + + +class TestRenderCoverLetter: + def test_returns_valid_pdf_bytes(self) -> None: + result = render_cover_letter("Dear Hiring Manager,\n\nI am applying...", SAMPLE_PROFILE) + assert bytes(result).startswith(b"%PDF") + + def test_swedish_text(self) -> None: + text = "Hej! Jag söker jobbet. Mvh Åke Öberg." + result = render_cover_letter(text, SWEDISH_PROFILE) + assert bytes(result).startswith(b"%PDF") + + def test_content_hash_stable(self) -> None: + text = "Cover letter body text." + result_a = render_cover_letter(text, SAMPLE_PROFILE) + result_b = render_cover_letter(text, SAMPLE_PROFILE) + assert hash_bytes(bytes(result_a)) == hash_bytes(bytes(result_b)) + + +class TestHashBytes: + def test_known_value(self) -> None: + assert hash_bytes(b"hello") == ( + "2cf24dba5fb0a30e26e83b2ac5b9e29e1b161e5c1fa7425e73043362938b9824" + ) + + def test_empty_bytes(self) -> None: + assert hash_bytes(b"") == ( + "e3b0c44298fc1c149afbf4c8996fb92427ae41e4649b934ca495991b7852b855" + ) + + def test_different_input_different_hash(self) -> None: + assert hash_bytes(b"a") != hash_bytes(b"b") + + +class TestNextVersion: + def test_empty_list(self) -> None: + assert next_version([]) == 1 + + def test_sequential(self) -> None: + assert next_version([1, 2, 3]) == 4 + + def test_gaps(self) -> None: + assert next_version([1, 3]) == 4 + + def test_single(self) -> None: + assert next_version([5]) == 6 \ No newline at end of file diff --git a/packages/llm-gateway/.gitkeep b/packages/llm-gateway/.gitkeep deleted file mode 100644 index e69de29..0000000 diff --git a/packages/llm-gateway/pyproject.toml b/packages/llm-gateway/pyproject.toml new file mode 100644 index 0000000..5a635a9 --- /dev/null +++ b/packages/llm-gateway/pyproject.toml @@ -0,0 +1,27 @@ +[project] +name = "llm-gateway" +version = "0.1.0" +description = "Async LLM gateway with mock mode, per-task budgets, telemetry, and retry/fallback." +requires-python = ">=3.13" +dependencies = [ + "httpx>=0.27", + "jsonschema>=4.23", +] + +[project.optional-dependencies] +dev = [ + "pytest>=8.0", + "pytest-asyncio>=0.24", +] + +[build-system] +requires = ["hatchling"] +build-backend = "hatchling.build" + +[tool.hatch.build.targets.wheel] +packages = ["src/llm_gateway"] + +[tool.pytest.ini_options] +testpaths = ["tests"] +pythonpath = ["src"] +asyncio_mode = "auto" \ No newline at end of file diff --git a/packages/llm-gateway/src/llm_gateway/__init__.py b/packages/llm-gateway/src/llm_gateway/__init__.py new file mode 100644 index 0000000..40d6df2 --- /dev/null +++ b/packages/llm-gateway/src/llm_gateway/__init__.py @@ -0,0 +1,27 @@ +"""LLM Gateway package. + +Async-first client with provider config from env, mock mode when no key is set, +per-task token budgets, telemetry sink, and retry/fallback policy. + +Public API: + Gateway: main async gateway class. + GatewayConfig: provider + budget configuration. + TelemetryRow: telemetry record dataclass. + BudgetExceeded: raised when a task would exceed its token budget. + run_task: convenience function using default config. +""" + +from llm_gateway.config import GatewayConfig, TaskClass +from llm_gateway.exceptions import BudgetExceeded, GatewayError, SchemaValidationError +from llm_gateway.gateway import Gateway, TelemetryRow, run_task + +__all__ = [ + "Gateway", + "GatewayConfig", + "TaskClass", + "TelemetryRow", + "BudgetExceeded", + "GatewayError", + "SchemaValidationError", + "run_task", +] \ No newline at end of file diff --git a/packages/llm-gateway/src/llm_gateway/config.py b/packages/llm-gateway/src/llm_gateway/config.py new file mode 100644 index 0000000..f12757e --- /dev/null +++ b/packages/llm-gateway/src/llm_gateway/config.py @@ -0,0 +1,218 @@ +"""Configuration for the LLM gateway. + +Provider config from env. Cheap task classes must never fall back to a paid +provider. Budgets are per-task max output tokens, configurable via env. +""" + +from __future__ import annotations + +import enum +import logging +import os +from dataclasses import dataclass, field + +logger = logging.getLogger(__name__) + + +class TaskClass(str, enum.Enum): + """Task classification for model routing. + + CHEAP tasks (score, extract) use the cheap model and must never fall back + to a paid provider. STRONG tasks (critique, prose review) may use fallback. + """ + + CHEAP = "cheap" + STRONG = "strong" + + +# Map task names to task classes for routing. +TASK_CLASS_MAP: dict[str, TaskClass] = { + "score": TaskClass.CHEAP, + "extract": TaskClass.CHEAP, + "cv_assist": TaskClass.CHEAP, + "cl_critique": TaskClass.STRONG, + "critique": TaskClass.STRONG, + "research": TaskClass.STRONG, +} + +# Default budgets (max output tokens) per task name. +DEFAULT_BUDGETS: dict[str, int] = { + "score": 2000, + "extract": 4000, + "cv_assist": 2000, + "cl_critique": 4000, + "critique": 6000, + "research": 4000, +} + + +@dataclass +class ProviderConfig: + """Configuration for a single LLM provider.""" + + name: str + base_url: str + api_key: str + model: str + is_paid: bool = False + + @property + def is_configured(self) -> bool: + """True if this provider has a non-empty API key.""" + return bool(self.api_key and self.api_key.strip()) + + +@dataclass +class GatewayConfig: + """Full gateway configuration loaded from environment. + + Attributes: + primary: primary provider config. + fallback: optional fallback provider config (None if not configured). + cheap_model: model name for cheap task classes. + strong_model: model name for strong task classes. + budgets: dict mapping task name to max output tokens. + max_retries: max retries on 429/5xx before fallback. + """ + + primary: ProviderConfig + fallback: ProviderConfig | None = None + cheap_model: str = "glm-5.2" + strong_model: str = "glm-5.2" + budgets: dict[str, int] = field(default_factory=lambda: dict(DEFAULT_BUDGETS)) + max_retries: int = 2 + + @property + def mock_mode(self) -> bool: + """True when no primary provider key is configured.""" + return not self.primary.is_configured + + @classmethod + def from_env(cls, env: dict[str, str] | None = None) -> GatewayConfig: + """Load configuration from environment variables. + + Env vars: + LLM_PRIMARY_BASE_URL, LLM_PRIMARY_KEY (or OLLAMA_API_KEY), + LLM_PRIMARY_MODEL (default glm-5.2) + LLM_FALLBACK_BASE_URL, LLM_FALLBACK_KEY, LLM_FALLBACK_MODEL + LLM_CHEAP_MODEL (default glm-5.2) + LLM_STRONG_MODEL (default glm-5.2) + LLM_BUDGET_{TASK} (per-task budget overrides) + LLM_MAX_RETRIES (default 2) + + A warning is logged (and in strict mode, ValueError raised) if a + paid fallback provider is configured while cheap task classes would + use it. The fallback is only used for STRONG tasks. + """ + e = env if env is not None else os.environ + + primary_key = e.get("LLM_PRIMARY_KEY", "") or e.get("OLLAMA_API_KEY", "") + primary = ProviderConfig( + name="primary", + base_url=e.get("LLM_PRIMARY_BASE_URL", "https://api.ollama-cloud.com/v1"), + api_key=primary_key, + model=e.get("LLM_PRIMARY_MODEL", "glm-5.2"), + is_paid=False, + ) + + fallback: ProviderConfig | None = None + fb_key = e.get("LLM_FALLBACK_KEY", "") + fb_url = e.get("LLM_FALLBACK_BASE_URL", "") + if fb_key and fb_url: + fallback = ProviderConfig( + name="fallback", + base_url=fb_url, + api_key=fb_key, + model=e.get("LLM_FALLBACK_MODEL", "glm-5.2"), + # Heuristic: if the base URL contains known paid provider hints, + # mark as paid. + is_paid=_detect_paid_provider(fb_url), + ) + if fallback.is_paid: + logger.warning( + "Fallback provider appears to be a paid provider (%s). " + "Cheap task classes will NOT use this fallback.", + fb_url, + ) + + cheap_model = e.get("LLM_CHEAP_MODEL", "glm-5.2") + strong_model = e.get("LLM_STRONG_MODEL", "glm-5.2") + + budgets = dict(DEFAULT_BUDGETS) + for task_name in list(budgets.keys()): + env_val = e.get(f"LLM_BUDGET_{task_name.upper()}") + if env_val: + try: + budgets[task_name] = int(env_val) + except ValueError: + logger.warning("Invalid budget value for %s: %s", task_name, env_val) + + # Also pick up any LLM_BUDGET_* not in defaults. + for key, val in e.items(): + if key.startswith("LLM_BUDGET_") and val: + task = key[len("LLM_BUDGET_"):].lower() + if task not in budgets: + try: + budgets[task] = int(val) + except ValueError: + pass + + max_retries_str = e.get("LLM_MAX_RETRIES", "2") + try: + max_retries = int(max_retries_str) + except ValueError: + max_retries = 2 + + return cls( + primary=primary, + fallback=fallback, + cheap_model=cheap_model, + strong_model=strong_model, + budgets=budgets, + max_retries=max_retries, + ) + + def get_budget(self, task: str) -> int: + """Return the max output token budget for *task*. + + Falls back to LLM_BUDGET_DEFAULT or 4000 if task-specific budget + is not set. + """ + return self.budgets.get(task, self.budgets.get("default", 4000)) + + def get_model(self, task: str) -> str: + """Return the model name for *task* based on its task class.""" + task_class = TASK_CLASS_MAP.get(task, TaskClass.STRONG) + if task_class == TaskClass.CHEAP: + return self.cheap_model + return self.strong_model + + def get_task_class(self, task: str) -> TaskClass: + """Return the task class for *task*.""" + return TASK_CLASS_MAP.get(task, TaskClass.STRONG) + + def assert_no_paid_fallback_for_cheap(self) -> None: + """Assert that no paid fallback is configured for cheap task classes. + + This is called during config validation. If a paid fallback exists, + it is allowed for STRONG tasks but must never be used for CHEAP tasks. + The gateway enforces this in _select_provider, but we also check here. + """ + if self.fallback and self.fallback.is_paid: + # This is allowed as long as cheap tasks never use fallback. + # We log a warning; the gateway itself prevents the routing. + logger.info( + "Paid fallback configured but will not be used for cheap tasks." + ) + + +def _detect_paid_provider(base_url: str) -> bool: + """Heuristic: detect if a base URL points to a known paid provider.""" + url_lower = base_url.lower() + paid_hints = [ + "openai.com", + "anthropic.com", + "api.openai.com", + "api.anthropic.com", + ] + return any(hint in url_lower for hint in paid_hints) \ No newline at end of file diff --git a/packages/llm-gateway/src/llm_gateway/exceptions.py b/packages/llm-gateway/src/llm_gateway/exceptions.py new file mode 100644 index 0000000..abf2fcc --- /dev/null +++ b/packages/llm-gateway/src/llm_gateway/exceptions.py @@ -0,0 +1,22 @@ +"""Exceptions for the LLM gateway.""" + +from __future__ import annotations + + +class GatewayError(Exception): + """Base exception for LLM gateway errors.""" + + +class BudgetExceeded(GatewayError): + """Raised when a task would exceed its configured token budget. + + This is raised BEFORE any provider call is made. + """ + + +class SchemaValidationError(GatewayError): + """Raised when the LLM output does not validate against the schema.""" + + +class ProviderError(GatewayError): + """Raised when a provider call fails after all retries and fallbacks.""" \ No newline at end of file diff --git a/packages/llm-gateway/src/llm_gateway/gateway.py b/packages/llm-gateway/src/llm_gateway/gateway.py new file mode 100644 index 0000000..27b6dd5 --- /dev/null +++ b/packages/llm-gateway/src/llm_gateway/gateway.py @@ -0,0 +1,388 @@ +"""Main gateway module: async-first LLM client with retry, fallback, budget guard, and telemetry.""" + +from __future__ import annotations + +import asyncio +import json +import logging +import time +import uuid +from dataclasses import dataclass, field +from typing import Any, Awaitable, Callable + +import httpx +from jsonschema import validate as jsonschema_validate +from jsonschema.exceptions import ValidationError as JsonSchemaValidationError + +from llm_gateway.config import GatewayConfig, TaskClass +from llm_gateway.exceptions import ( + BudgetExceeded, + GatewayError, + ProviderError, + SchemaValidationError, +) +from llm_gateway.mock import get_mock_output, mock_telemetry_row + +logger = logging.getLogger(__name__) + +# Type alias for the telemetry sink: an async callable that receives a +# TelemetryRow (or a dict for simple sinks). +TelemetrySink = Callable[["TelemetryRow"], Awaitable[None] | None] + + +@dataclass +class TelemetryRow: + """A single telemetry record for one LLM call. + + Attributes: + id: unique run id. + task: task name (e.g. 'score', 'extract'). + model: model name used. + provider: provider name ('primary', 'fallback', 'mock'). + input_tokens: tokens consumed on input. + output_tokens: tokens consumed on output. + cost_usd: estimated cost in USD (None if not configured). + duration_ms: wall-clock duration in milliseconds. + application_id: optional application id for correlation. + mock: True if this was a mock-mode call. + """ + + id: str = field(default_factory=lambda: str(uuid.uuid4())) + task: str = "" + model: str = "" + provider: str = "" + input_tokens: int = 0 + output_tokens: int = 0 + cost_usd: float | None = None + duration_ms: int = 0 + application_id: str | None = None + mock: bool = False + + def to_dict(self) -> dict[str, Any]: + """Serialize to dict for sinks that accept plain dicts.""" + return { + "id": self.id, + "task": self.task, + "model": self.model, + "provider": self.provider, + "input_tokens": self.input_tokens, + "output_tokens": self.output_tokens, + "cost_usd": self.cost_usd, + "duration_ms": self.duration_ms, + "application_id": self.application_id, + "mock": self.mock, + } + + +class Gateway: + """Async-first LLM gateway. + + Usage: + config = GatewayConfig.from_env() + gw = Gateway(config) + result = await gw.run_task("score", "Score this job vs profile: ...") + + When no API key is configured (mock mode), returns deterministic canned + outputs without making any network calls. + """ + + def __init__( + self, + config: GatewayConfig, + telemetry_sink: TelemetrySink | None = None, + http_client: httpx.AsyncClient | None = None, + ) -> None: + self.config = config + self._telemetry_sink = telemetry_sink + self._http_client = http_client + self._owns_http_client = http_client is None + + async def _get_http_client(self) -> httpx.AsyncClient: + if self._http_client is None: + self._http_client = httpx.AsyncClient(timeout=30.0) + return self._http_client + + async def aclose(self) -> None: + """Close the HTTP client if we own it.""" + if self._owns_http_client and self._http_client is not None: + await self._http_client.aclose() + self._http_client = None + + async def __aenter__(self) -> Gateway: + return self + + async def __aexit__(self, *args: Any) -> None: + await self.aclose() + + def _select_provider(self, task: str) -> Any: + """Select the provider config for *task*. + + Cheap tasks always use the primary provider (never fallback to paid). + Strong tasks may use fallback if primary fails. + """ + task_class = self.config.get_task_class(task) + # Cheap tasks: primary only, never fallback (especially not paid). + if task_class == TaskClass.CHEAP: + return self.config.primary + # Strong tasks: primary, with fallback available. + return self.config.primary + + def _check_budget(self, task: str, prompt: str) -> None: + """Check if the estimated token usage would exceed the budget. + + Raises BudgetExceeded BEFORE any provider call is made. + We estimate input tokens as len(prompt) // 4 (rough heuristic) and + add the max output token budget. If the estimated total exceeds + the budget, we raise. + """ + budget = self.config.get_budget(task) + # Rough input token estimate: ~4 chars per token. + estimated_input_tokens = len(prompt) // 4 + # If input alone exceeds budget, that is over-budget. + if estimated_input_tokens > budget: + raise BudgetExceeded( + f"Task '{task}' estimated input tokens ({estimated_input_tokens}) " + f"exceed budget ({budget}). Call aborted before provider request." + ) + + async def _write_telemetry(self, row: TelemetryRow) -> None: + """Send telemetry to the sink if one is configured.""" + if self._telemetry_sink is None: + return + result = self._telemetry_sink(row) + if asyncio.iscoroutine(result): + await result + + async def _call_provider( + self, + provider_config: Any, + task: str, + prompt: str, + model: str, + ) -> dict[str, Any]: + """Make an async HTTP call to the provider's chat completions endpoint. + + Returns the raw JSON response dict. + """ + client = await self._get_http_client() + url = f"{provider_config.base_url.rstrip('/')}/chat/completions" + headers = { + "Authorization": f"Bearer {provider_config.api_key}", + "Content-Type": "application/json", + } + body: dict[str, Any] = { + "model": model, + "messages": [{"role": "user", "content": prompt}], + "max_tokens": self.config.get_budget(task), + } + # If a schema is expected, request JSON format. + body["response_format"] = {"type": "json_object"} + + response = await client.post(url, json=body, headers=headers) + response.raise_for_status() + data = response.json() + + # Extract the content from the response. + choices = data.get("choices", []) + if not choices: + raise ProviderError(f"Provider returned no choices for task '{task}'") + content = choices[0].get("message", {}).get("content", "{}") + usage = data.get("usage", {}) + + try: + parsed = json.loads(content) + except json.JSONDecodeError as exc: + raise ProviderError( + f"Provider returned non-JSON content for task '{task}': {exc}" + ) from exc + + # Attach usage info for telemetry. + parsed["_usage"] = { + "input_tokens": usage.get("prompt_tokens", 0), + "output_tokens": usage.get("completion_tokens", 0), + "model": data.get("model", model), + "provider": provider_config.name, + } + return parsed + + async def _call_with_retry( + self, + provider_config: Any, + task: str, + prompt: str, + model: str, + ) -> dict[str, Any]: + """Call provider with retry policy: max retries on 429/5xx, then + fallback provider (for strong tasks only), then raise. + """ + last_exc: Exception | None = None + for attempt in range(self.config.max_retries + 1): + try: + return await self._call_provider(provider_config, task, prompt, model) + except httpx.HTTPStatusError as exc: + status = exc.response.status_code + if status == 429 or 500 <= status < 600: + last_exc = exc + logger.warning( + "Provider %s returned %d for task '%s' (attempt %d/%d)", + provider_config.name, + status, + task, + attempt + 1, + self.config.max_retries + 1, + ) + if attempt < self.config.max_retries: + await asyncio.sleep(2 ** attempt) + continue + raise + except (httpx.RequestError, ProviderError) as exc: + last_exc = exc + logger.warning( + "Provider %s error for task '%s' (attempt %d/%d): %s", + provider_config.name, + task, + attempt + 1, + self.config.max_retries + 1, + exc, + ) + if attempt < self.config.max_retries: + await asyncio.sleep(2 ** attempt) + continue + + # All retries exhausted. Try fallback for strong tasks. + task_class = self.config.get_task_class(task) + if ( + task_class == TaskClass.STRONG + and self.config.fallback + and self.config.fallback.is_configured + and self.config.fallback is not provider_config + ): + logger.info("Falling back to %s for task '%s'", self.config.fallback.name, task) + try: + return await self._call_provider( + self.config.fallback, task, prompt, self.config.strong_model + ) + except Exception as fallback_exc: + raise ProviderError( + f"Both primary and fallback providers failed for task '{task}': " + f"primary={last_exc}, fallback={fallback_exc}" + ) from fallback_exc + + raise ProviderError( + f"Provider call failed for task '{task}' after {self.config.max_retries + 1} attempts: {last_exc}" + ) from last_exc + + async def run_task( + self, + task: str, + prompt: str, + schema: dict | None = None, + application_id: str | None = None, + ) -> dict[str, Any]: + """Run an LLM task and return the parsed result. + + Args: + task: task name (e.g. 'score', 'extract', 'cv_assist', 'cl_critique'). + prompt: the input prompt text. + schema: optional JSON schema to validate the output against. + application_id: optional application id for telemetry correlation. + + Returns: + Parsed JSON dict from the LLM. + + Raises: + BudgetExceeded: if the estimated token usage exceeds the budget. + SchemaValidationError: if the output does not match the schema. + ProviderError: if the provider call fails after all retries. + """ + # Budget guard: raise BEFORE any call is made. + self._check_budget(task, prompt) + + start = time.monotonic() + + if self.config.mock_mode: + # Mock mode: return deterministic canned output. + result = get_mock_output(task) + duration_ms = int((time.monotonic() - start) * 1000) + model = self.config.get_model(task) + + row = TelemetryRow( + task=task, + model=model, + provider="mock", + input_tokens=len(prompt) // 4, + output_tokens=0, + cost_usd=0.0, + duration_ms=duration_ms, + application_id=application_id, + mock=True, + ) + await self._write_telemetry(row) + + # Validate against schema if provided. + if schema is not None: + _validate_schema(result, schema) + + return result + + # Real mode: call provider with retry/fallback. + provider = self._select_provider(task) + model = self.config.get_model(task) + + try: + result = await self._call_with_retry(provider, task, prompt, model) + except BudgetExceeded: + raise + except (ProviderError, SchemaValidationError): + raise + except Exception as exc: + raise ProviderError(f"Unexpected error for task '{task}': {exc}") from exc + + duration_ms = int((time.monotonic() - start) * 1000) + usage = result.pop("_usage", {}) + + row = TelemetryRow( + task=task, + model=usage.get("model", model), + provider=usage.get("provider", provider.name), + input_tokens=usage.get("input_tokens", 0), + output_tokens=usage.get("output_tokens", 0), + cost_usd=None, + duration_ms=duration_ms, + application_id=application_id, + mock=False, + ) + await self._write_telemetry(row) + + # Validate against schema if provided. + if schema is not None: + _validate_schema(result, schema) + + return result + + +def _validate_schema(data: dict, schema: dict) -> None: + """Validate *data* against *schema*. Raises SchemaValidationError on failure.""" + try: + jsonschema_validate(data, schema) + except JsonSchemaValidationError as exc: + raise SchemaValidationError(f"Schema validation failed: {exc.message}") from exc + + +async def run_task( + task: str, + prompt: str, + schema: dict | None = None, + config: GatewayConfig | None = None, + telemetry_sink: TelemetrySink | None = None, +) -> dict[str, Any]: + """Convenience function: create a Gateway, run one task, close it. + + Useful for one-off calls. For repeated calls, instantiate Gateway directly. + """ + cfg = config or GatewayConfig.from_env() + gw = Gateway(cfg, telemetry_sink=telemetry_sink) + try: + return await gw.run_task(task, prompt, schema=schema) + finally: + await gw.aclose() \ No newline at end of file diff --git a/packages/llm-gateway/src/llm_gateway/mock.py b/packages/llm-gateway/src/llm_gateway/mock.py new file mode 100644 index 0000000..598bfa8 --- /dev/null +++ b/packages/llm-gateway/src/llm_gateway/mock.py @@ -0,0 +1,96 @@ +"""Mock mode for the LLM gateway. + +When no API key is configured, the gateway returns deterministic canned +outputs per task name. This allows the API and tests to run offline. +""" + +from __future__ import annotations + +import json +import time +import uuid + +# Deterministic canned outputs per task name. +# Each entry is a dict that will be returned as the task result. +MOCK_OUTPUTS: dict[str, dict] = { + "score": { + "score": 75, + "rationale": { + "match": "good", + "reasons": ["skills align", "location matches"], + }, + }, + "extract": { + "company": "Example Corp", + "title": "Software Engineer", + "location": "Stockholm", + "requirements": ["Python", "PostgreSQL", "Docker"], + }, + "cv_assist": { + "suggestions": [ + "Led a team of 5 developers to deliver a critical integration", + "Reduced API latency by 40% through caching and query optimization", + ], + }, + "cl_critique": { + "comments": [ + { + "quote": "I am a hard worker", + "suggestion": "Replace generic claim with a specific achievement metric", + "severity": "medium", + }, + { + "quote": "Dear Sir/Madam", + "suggestion": "Address the hiring manager by name if known", + "severity": "low", + }, + ], + }, + "critique": { + "comments": [ + { + "quote": "sample text", + "suggestion": "improve clarity", + "severity": "low", + }, + ], + }, + "research": { + "summary": "The company is a mid-size tech firm focused on cloud infrastructure.", + "key_points": ["Founded in 2015", "Series B funding", "Remote-first culture"], + }, +} + +# Default mock output for unknown task names. +DEFAULT_MOCK_OUTPUT: dict = { + "result": "mock output", + "task": "unknown", +} + + +def get_mock_output(task: str) -> dict: + """Return a deterministic mock output for *task*. + + For unknown tasks, returns DEFAULT_MOCK_OUTPUT with the task name filled in. + """ + if task in MOCK_OUTPUTS: + # Return a copy so callers cannot mutate the canned data. + return json.loads(json.dumps(MOCK_OUTPUTS[task])) + result = json.loads(json.dumps(DEFAULT_MOCK_OUTPUT)) + result["task"] = task + return result + + +def mock_telemetry_row(task: str, model: str) -> dict: + """Build a mock telemetry row dict for offline mode.""" + return { + "id": str(uuid.uuid4()), + "task": task, + "model": model, + "provider": "mock", + "input_tokens": 0, + "output_tokens": 0, + "cost_usd": 0.0, + "duration_ms": 0, + "mock": True, + } \ No newline at end of file diff --git a/packages/llm-gateway/tests/test_gateway.py b/packages/llm-gateway/tests/test_gateway.py new file mode 100644 index 0000000..04786ad --- /dev/null +++ b/packages/llm-gateway/tests/test_gateway.py @@ -0,0 +1,491 @@ +"""Tests for the LLM gateway package.""" + +from __future__ import annotations + +import asyncio +import json +from typing import Any + +import httpx +import pytest + +from llm_gateway.config import GatewayConfig, ProviderConfig, TaskClass +from llm_gateway.exceptions import BudgetExceeded, SchemaValidationError +from llm_gateway.gateway import Gateway, TelemetryRow, run_task +from llm_gateway.mock import get_mock_output, MOCK_OUTPUTS + + +# ---- Fixtures ---- + + +def mock_config(**overrides: Any) -> GatewayConfig: + """Build a config in mock mode (no API key).""" + primary = ProviderConfig( + name="primary", + base_url="https://mock.example.com/v1", + api_key="", + model="glm-5.2", + ) + defaults: dict[str, Any] = { + "primary": primary, + "fallback": None, + "cheap_model": "glm-5.2", + "strong_model": "glm-5.2", + "budgets": {"score": 2000, "extract": 4000, "default": 4000}, + "max_retries": 2, + } + defaults.update(overrides) + return GatewayConfig(**defaults) + + +def real_config(**overrides: Any) -> GatewayConfig: + """Build a config with a fake API key (non-mock mode).""" + primary = ProviderConfig( + name="primary", + base_url="https://mock.example.com/v1", + api_key="fake-key-1234", + model="glm-5.2", + ) + defaults: dict[str, Any] = { + "primary": primary, + "fallback": None, + "cheap_model": "glm-5.2", + "strong_model": "glm-5.2", + "budgets": {"score": 2000, "extract": 4000, "default": 4000}, + "max_retries": 2, + } + defaults.update(overrides) + return GatewayConfig(**defaults) + + +# ---- Mock mode tests ---- + + +class TestMockMode: + async def test_mock_returns_deterministic_output(self) -> None: + """Mock mode returns deterministic canned outputs per task.""" + config = mock_config() + gw = Gateway(config) + result_a = await gw.run_task("score", "Score this job") + result_b = await gw.run_task("score", "Score this job") + assert result_a == result_b + assert result_a["score"] == 75 + await gw.aclose() + + async def test_mock_different_tasks_different_output(self) -> None: + config = mock_config() + gw = Gateway(config) + score_result = await gw.run_task("score", "prompt") + extract_result = await gw.run_task("extract", "prompt") + assert score_result != extract_result + assert "score" in score_result + assert "company" in extract_result + await gw.aclose() + + async def test_mock_unknown_task(self) -> None: + config = mock_config() + gw = Gateway(config) + result = await gw.run_task("unknown_task", "prompt") + assert result["result"] == "mock output" + assert result["task"] == "unknown_task" + await gw.aclose() + + async def test_mock_mode_property(self) -> None: + config = mock_config() + assert config.mock_mode is True + + async def test_real_mode_not_mock(self) -> None: + config = real_config() + assert config.mock_mode is False + + +# ---- Schema validation tests ---- + + +class TestSchemaValidation: + async def test_schema_passes(self) -> None: + config = mock_config() + gw = Gateway(config) + schema = { + "type": "object", + "properties": { + "score": {"type": "number"}, + "rationale": {"type": "object"}, + }, + "required": ["score"], + } + result = await gw.run_task("score", "prompt", schema=schema) + assert "score" in result + await gw.aclose() + + async def test_schema_fails(self) -> None: + """Schema validation failure should raise SchemaValidationError.""" + config = mock_config() + gw = Gateway(config) + # The mock output for 'score' has score=75 (number). We require a string, + # which should fail validation. + bad_schema = { + "type": "object", + "properties": { + "score": {"type": "string"}, + }, + "required": ["score"], + } + with pytest.raises(SchemaValidationError): + await gw.run_task("score", "prompt", schema=bad_schema) + await gw.aclose() + + async def test_schema_missing_required_field(self) -> None: + config = mock_config() + gw = Gateway(config) + schema = { + "type": "object", + "required": ["nonexistent_field"], + } + with pytest.raises(SchemaValidationError): + await gw.run_task("score", "prompt", schema=schema) + await gw.aclose() + + +# ---- Budget guard tests ---- + + +class TestBudgetGuard: + async def test_budget_exceeded_raises_before_call(self) -> None: + """Over-budget prompt should raise BudgetExceeded before any call.""" + config = mock_config(budgets={"score": 10}) + gw = Gateway(config) + # 10 token budget, ~4 chars/token, so >40 chars should exceed. + long_prompt = "x" * 100 + with pytest.raises(BudgetExceeded): + await gw.run_task("score", long_prompt) + await gw.aclose() + + async def test_budget_within_limit_does_not_raise(self) -> None: + config = mock_config(budgets={"score": 10000}) + gw = Gateway(config) + result = await gw.run_task("score", "short prompt") + assert result["score"] == 75 + await gw.aclose() + + async def test_budget_guard_in_real_mode(self) -> None: + """Budget guard must raise before call even in real (non-mock) mode.""" + config = real_config(budgets={"score": 10}) + gw = Gateway(config) + with pytest.raises(BudgetExceeded): + await gw.run_task("score", "x" * 100) + await gw.aclose() + + async def test_default_budget_fallback(self) -> None: + """Unknown task should use default budget.""" + config = mock_config(budgets={"score": 2000, "default": 100}) + gw = Gateway(config) + # Unknown task uses default=100, so >400 chars exceeds. + with pytest.raises(BudgetExceeded): + await gw.run_task("unknown_task", "x" * 500) + await gw.aclose() + + +# ---- Telemetry tests ---- + + +class TestTelemetry: + async def test_telemetry_sink_called_in_mock_mode(self) -> None: + sink_calls: list[TelemetryRow] = [] + + async def sink(row: TelemetryRow) -> None: + sink_calls.append(row) + + config = mock_config() + gw = Gateway(config, telemetry_sink=sink) + await gw.run_task("score", "test prompt") + assert len(sink_calls) == 1 + assert sink_calls[0].task == "score" + assert sink_calls[0].mock is True + await gw.aclose() + + async def test_telemetry_sink_sync_callable(self) -> None: + """Sync sinks should also work (no await needed).""" + sink_calls: list[TelemetryRow] = [] + + def sync_sink(row: TelemetryRow) -> None: + sink_calls.append(row) + + config = mock_config() + gw = Gateway(config, telemetry_sink=sync_sink) + await gw.run_task("score", "test prompt") + assert len(sink_calls) == 1 + await gw.aclose() + + async def test_no_sink_no_error(self) -> None: + config = mock_config() + gw = Gateway(config, telemetry_sink=None) + result = await gw.run_task("score", "prompt") + assert result["score"] == 75 + await gw.aclose() + + async def test_telemetry_row_to_dict(self) -> None: + row = TelemetryRow(task="score", model="glm-5.2", provider="mock") + d = row.to_dict() + assert d["task"] == "score" + assert d["model"] == "glm-5.2" + assert d["provider"] == "mock" + assert "id" in d + + +# ---- Config tests ---- + + +class TestGatewayConfig: + def test_from_env_mock_mode(self) -> None: + env = {} + config = GatewayConfig.from_env(env=env) + assert config.mock_mode is True + assert config.primary.model == "glm-5.2" + + def test_from_env_real_mode(self) -> None: + env = { + "LLM_PRIMARY_KEY": "test-key", + "LLM_PRIMARY_BASE_URL": "https://api.example.com/v1", + "LLM_PRIMARY_MODEL": "custom-model", + } + config = GatewayConfig.from_env(env=env) + assert config.mock_mode is False + assert config.primary.api_key == "test-key" + assert config.primary.model == "custom-model" + + def test_from_env_ollama_key(self) -> None: + env = {"OLLAMA_API_KEY": "ollama-key-123"} + config = GatewayConfig.from_env(env=env) + assert config.mock_mode is False + assert config.primary.api_key == "ollama-key-123" + + def test_from_env_budgets(self) -> None: + env = {"LLM_BUDGET_SCORE": "500"} + config = GatewayConfig.from_env(env=env) + assert config.budgets["score"] == 500 + + def test_task_class_routing(self) -> None: + config = mock_config() + assert config.get_task_class("score") == TaskClass.CHEAP + assert config.get_task_class("extract") == TaskClass.CHEAP + assert config.get_task_class("cv_assist") == TaskClass.CHEAP + assert config.get_task_class("critique") == TaskClass.STRONG + assert config.get_task_class("cl_critique") == TaskClass.STRONG + + def test_get_model_routing(self) -> None: + config = mock_config(cheap_model="cheap-model", strong_model="strong-model") + assert config.get_model("score") == "cheap-model" + assert config.get_model("critique") == "strong-model" + + def test_paid_fallback_detection(self) -> None: + """Config should detect paid provider URLs.""" + env = { + "LLM_PRIMARY_KEY": "key", + "LLM_FALLBACK_BASE_URL": "https://api.openai.com/v1", + "LLM_FALLBACK_KEY": "fb-key", + } + config = GatewayConfig.from_env(env=env) + assert config.fallback is not None + assert config.fallback.is_paid is True + + def test_paid_fallback_not_used_for_cheap(self) -> None: + """The gateway must not route cheap tasks to paid fallback. + + We verify this by checking _select_provider returns primary for cheap. + """ + primary = ProviderConfig( + name="primary", base_url="https://a.com/v1", api_key="k", model="m" + ) + fallback = ProviderConfig( + name="fallback", + base_url="https://api.openai.com/v1", + api_key="k2", + model="m2", + is_paid=True, + ) + config = GatewayConfig(primary=primary, fallback=fallback) + gw = Gateway(config) + provider = gw._select_provider("score") + assert provider.name == "primary" + # The _select_provider method enforces this by always returning primary + # for cheap tasks, never the paid fallback. + + def test_non_paid_fallback(self) -> None: + """Non-paid (e.g. ollama) fallback is fine.""" + env = { + "LLM_PRIMARY_KEY": "key", + "LLM_FALLBACK_BASE_URL": "https://api.ollama-cloud.com/v1", + "LLM_FALLBACK_KEY": "fb-key", + } + config = GatewayConfig.from_env(env=env) + assert config.fallback is not None + assert config.fallback.is_paid is False + + +# ---- Provider call tests (with mocked HTTP) ---- + + +class TestProviderCalls: + async def test_real_mode_calls_provider(self) -> None: + """In non-mock mode, the gateway should make an HTTP call.""" + config = real_config() + mock_response_data = { + "choices": [ + { + "message": { + "content": json.dumps({"score": 85, "rationale": {"ok": True}}) + } + } + ], + "usage": {"prompt_tokens": 50, "completion_tokens": 30}, + "model": "glm-5.2", + } + + def handler(request: httpx.Request) -> httpx.Response: + return httpx.Response(200, json=mock_response_data) + + client = httpx.AsyncClient( + transport=httpx.MockTransport(handler), + base_url="https://mock.example.com", + ) + gw = Gateway(config, http_client=client) + result = await gw.run_task("score", "Score this job") + assert result["score"] == 85 + assert "_usage" not in result # _usage should be popped + await client.aclose() + + async def test_retry_on_429(self) -> None: + """Gateway should retry on 429 then succeed.""" + config = real_config(max_retries=2) + call_count = 0 + + def handler(request: httpx.Request) -> httpx.Response: + nonlocal call_count + call_count += 1 + if call_count < 2: + return httpx.Response(429, json={"error": "rate limited"}) + return httpx.Response( + 200, + json={ + "choices": [ + {"message": {"content": json.dumps({"score": 50})}} + ], + "usage": {"prompt_tokens": 10, "completion_tokens": 5}, + }, + ) + + client = httpx.AsyncClient( + transport=httpx.MockTransport(handler), + base_url="https://mock.example.com", + ) + gw = Gateway(config, http_client=client) + result = await gw.run_task("score", "test") + assert result["score"] == 50 + assert call_count == 2 + await client.aclose() + + async def test_retry_exhausted_raises(self) -> None: + """After all retries, ProviderError should be raised.""" + config = real_config(max_retries=1) + + def handler(request: httpx.Request) -> httpx.Response: + return httpx.Response(500, json={"error": "server error"}) + + client = httpx.AsyncClient( + transport=httpx.MockTransport(handler), + base_url="https://mock.example.com", + ) + from llm_gateway.exceptions import ProviderError + + gw = Gateway(config, http_client=client) + with pytest.raises(ProviderError): + await gw.run_task("critique", "test") + await client.aclose() + + async def test_fallback_used_for_strong_task(self) -> None: + """Strong tasks should fall back when primary fails.""" + primary = ProviderConfig( + name="primary", + base_url="https://primary.example.com/v1", + api_key="pk", + model="glm-5.2", + ) + fallback = ProviderConfig( + name="fallback", + base_url="https://fallback.example.com/v1", + api_key="fk", + model="glm-5.2", + is_paid=False, + ) + config = GatewayConfig( + primary=primary, + fallback=fallback, + budgets={"critique": 4000, "default": 4000}, + max_retries=0, + ) + + def handler(request: httpx.Request) -> httpx.Response: + if "primary.example.com" in str(request.url): + return httpx.Response(500, json={"error": "primary down"}) + return httpx.Response( + 200, + json={ + "choices": [ + {"message": {"content": json.dumps({"comments": []})}} + ], + "usage": {"prompt_tokens": 20, "completion_tokens": 10}, + "model": "glm-5.2", + }, + ) + + client = httpx.AsyncClient( + transport=httpx.MockTransport(handler), + ) + gw = Gateway(config, http_client=client) + result = await gw.run_task("critique", "review this") + assert result == {"comments": []} + await client.aclose() + + async def test_no_fallback_for_cheap_task(self) -> None: + """Cheap tasks must not use fallback even when primary fails.""" + primary = ProviderConfig( + name="primary", + base_url="https://primary.example.com/v1", + api_key="pk", + model="glm-5.2", + ) + fallback = ProviderConfig( + name="fallback", + base_url="https://fallback.example.com/v1", + api_key="fk", + model="glm-5.2", + ) + config = GatewayConfig( + primary=primary, + fallback=fallback, + budgets={"score": 4000, "default": 4000}, + max_retries=0, + ) + + def handler(request: httpx.Request) -> httpx.Response: + return httpx.Response(500, json={"error": "down"}) + + client = httpx.AsyncClient( + transport=httpx.MockTransport(handler), + ) + from llm_gateway.exceptions import ProviderError + + gw = Gateway(config, http_client=client) + with pytest.raises(ProviderError): + await gw.run_task("score", "score this") + await client.aclose() + + +# ---- Convenience function test ---- + + +class TestRunTaskFunction: + async def test_run_task_convenience_mock(self) -> None: + result = await run_task("score", "test", config=mock_config()) + assert result["score"] == 75 \ No newline at end of file