T2: packages/artifacts + packages/llm-gateway + docker-compose.yml + .env.example

packages/artifacts:
- render_cv_pdf(profile, sections) -> bytes: data-driven CV PDF generation
  using fpdf2 with bundled DejaVuSans TTF for unicode (Swedish chars tested).
  Jinja2 template for layout data prep, adapted from build_cv.py approach.
- render_cover_letter(text, profile) -> bytes: simple cover letter PDF.
- hash_bytes(b) -> str: sha256 hex digest.
- next_version(existing) -> int: version numbering helper.
- 16 tests, all passing: PDF validity, Swedish characters, hash stability,
  cover letter rendering, hash correctness, version logic.

packages/llm-gateway:
- Async-first Gateway class with provider config from env.
- Mock mode default when no API key env present (deterministic canned outputs
  per task name, defined in mock.py).
- Telemetry sink injectable (async or sync callable, receives TelemetryRow).
- Budget guard raises BudgetExceeded BEFORE any provider call is made.
- Retry policy: max 2 retries on 429/5xx, then fallback provider for STRONG
  tasks only. CHEAP tasks never use fallback (paid provider protection).
- Schema validation via jsonschema; SchemaValidationError on mismatch.
- Task class routing: CHEAP (score, extract, cv_assist) vs STRONG (critique,
  cl_critique, research). Model routing per task class.
- Paid provider detection heuristic; warns on paid fallback config.
- 31 tests, all passing: mock determinism, schema pass/fail, budget guard
  (mock + real mode), telemetry sink (async/sync/none), config from env,
  provider calls with mocked HTTP (retry, fallback, no-fallback-for-cheap).

docker-compose.yml:
- postgres:16 service, user/pass/db = jobhunt, host port 5433->5432,
  named volume jobhunt_pgdata, healthcheck.

.env.example:
- DATABASE_URL, LLM provider config (primary + fallback), task budgets,
  API and web settings.
This commit is contained in:
hermes 2026-07-30 17:59:56 +00:00
parent b77c8b0044
commit 8d8a863300
18 changed files with 1804 additions and 0 deletions

37
.env.example Normal file
View file

@ -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

23
docker-compose.yml Normal file
View file

@ -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

View file

@ -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"

Binary file not shown.

Binary file not shown.

View file

@ -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",
]

View file

@ -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()

View file

@ -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

View file

@ -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

View file

@ -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"

View file

@ -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",
]

View file

@ -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)

View file

@ -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."""

View file

@ -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()

View file

@ -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,
}

View file

@ -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