diff --git a/fleet/__init__.py b/fleet/__init__.py index 6a429dab..e2e649c4 100644 --- a/fleet/__init__.py +++ b/fleet/__init__.py @@ -74,6 +74,9 @@ # Import judge data classes from .judge import Rubric, Criterion, File, Image, JudgeResult +# Import LLM provider interface +from .llm_provider import LLMProvider, FleetProvider, ExternalProvider, resolve_provider + # Create a module-level env attribute for convenient access from . import env from . import global_client as _global_client @@ -102,6 +105,11 @@ "File", "Image", "JudgeResult", + # LLM Providers + "LLMProvider", + "FleetProvider", + "ExternalProvider", + "resolve_provider", # Exceptions "FleetError", "FleetAPIError", diff --git a/fleet/_async/client.py b/fleet/_async/client.py index 80013269..d0399fe3 100644 --- a/fleet/_async/client.py +++ b/fleet/_async/client.py @@ -55,6 +55,7 @@ if TYPE_CHECKING: from .verifiers import AsyncVerifierFunction from .judge import AsyncJudge + from ..llm_provider import LLMProvider def _json_default(x: Any) -> Any: @@ -340,12 +341,19 @@ def message_count(self) -> int: class AsyncEnv(EnvironmentBase): - def __init__(self, client: Optional[AsyncWrapper], **kwargs): + def __init__( + self, + client: Optional[AsyncWrapper], + *, + llm_provider: Optional["LLMProvider"] = None, + **kwargs, + ): super().__init__(**kwargs) self._client = client self._apps: Dict[str, AsyncInstanceClient] = {} self._instance: Optional[AsyncInstanceClient] = None self._judge: Optional["AsyncJudge"] = None + self._llm_provider = llm_provider @property def instance(self) -> AsyncInstanceClient: @@ -423,13 +431,18 @@ def mcp(self) -> AsyncMCPResource: @property def judge(self) -> "AsyncJudge": - """LLM-as-judge grading via orchestrator API.""" + """LLM-as-judge grading. + + Routes through Fleet orchestrator by default. Set ``llm_provider`` + on the environment to route to an external provider instead. + """ if self._judge is None: from .judge import AsyncJudge self._judge = AsyncJudge( - client=self._load_client, + client=self._client, instance_id=self.instance_id, + llm_provider=self._llm_provider, ) return self._judge diff --git a/fleet/_async/judge.py b/fleet/_async/judge.py index fe1d2bae..ede4e38b 100644 --- a/fleet/_async/judge.py +++ b/fleet/_async/judge.py @@ -1,9 +1,15 @@ """Fleet SDK Judge - Async version. Provides env.judge.grade() for async verifier scripts. + +Provider resolution order: + +1. Explicit ``llm_provider`` kwarg (highest priority) +2. ``FLEET_LLM_API_KEY`` env var → auto-builds ``ExternalProvider`` +3. Fleet orchestrator (default fallback) """ -from typing import Dict, List, Optional, Union, TYPE_CHECKING +from typing import Any, Dict, List, Optional, Union, TYPE_CHECKING # Import shared classes and helpers from the sync module from ..judge import ( @@ -19,10 +25,12 @@ _guess_media_type, _parse_grade_response, _print_judge_call_start, + _UNSET, ) if TYPE_CHECKING: from .base import AsyncWrapper + from ..llm_provider import LLMProvider # Re-export data classes so `from fleet._async.judge import ...` works __all__ = [ @@ -36,15 +44,33 @@ class AsyncJudge: - """LLM-as-judge grading — calls orchestrator API, not environment API. + """LLM-as-judge grading (async). + + Accessed as ``env.judge`` on AsyncEnv instances. - Accessed as env.judge on AsyncEnv instances. + Provider resolution order: + + 1. Explicit ``llm_provider`` kwarg (highest priority) + 2. ``FLEET_LLM_API_KEY`` env var → auto-builds ``ExternalProvider`` + 3. Fleet orchestrator (default fallback) """ - def __init__(self, client: "AsyncWrapper", instance_id: str): + def __init__( + self, + client: Optional["AsyncWrapper"], + instance_id: str, + *, + llm_provider: Any = _UNSET, + ): self._client = client self._instance_id = instance_id + if llm_provider is _UNSET: + from ..llm_provider import resolve_provider + self._llm_provider = resolve_provider() + else: + self._llm_provider = llm_provider + async def grade( self, rubric: Union[str, Rubric], @@ -63,7 +89,11 @@ async def grade( collect: Optional[Dict[str, List[str]]] = None, task_id: Optional[str] = None, ) -> JudgeResult: - """Grade a submission using LLM-as-judge via the orchestrator API. + """Grade a submission using LLM-as-judge. + + Routes through the Fleet orchestrator by default. If an + ``llm_provider`` was set at construction time, calls the external + provider directly instead. Returns a JudgeResult (float subclass with .details, .criteria, .feedback) that can be returned directly from a verifier function. @@ -84,15 +114,29 @@ async def grade( collect: File patterns for orchestrator to collect (agentic mode). task_id: Optional task ID for tracking. """ - # Resolve Image.from_env images asynchronously before building request + # Fold reference_claims into context + effective_context = context + if reference_claims is not None: + if effective_context: + effective_context = f"{effective_context}\n\n## Reference Claims\n{reference_claims}" + else: + effective_context = f"## Reference Claims\n{reference_claims}" + + # Resolve path-based images/files through the provider first resolved_images = images - if images and not agentic: - resolved_images = {} - for label, img in images.items(): + resolved_files = files + if self._llm_provider is not None: + resolved_images = self._llm_provider.resolve_images(images) + resolved_files = self._llm_provider.resolve_files(files) + + # Resolve Image.from_env images asynchronously before building request + if resolved_images and not agentic: + env_resolved = {} + for label, img in resolved_images.items(): if img.source == "env" and img._env is not None: b64 = await _collect_image_from_env_async(img._env, img.filename) if b64 is not None: - resolved_images[label] = Image.from_base64( + env_resolved[label] = Image.from_base64( b64, img.filename or "image.png", _guess_media_type(img.filename or "image.png"), @@ -100,43 +144,70 @@ async def grade( else: # Async collection failed — use collect source directly # (don't keep the env image or serialize() will retry sync) - resolved_images[label] = Image( + env_resolved[label] = Image( source="collect", filename=img.filename, ) else: - resolved_images[label] = img + env_resolved[label] = img + resolved_images = env_resolved # Resolve File.from_env files asynchronously before building request - resolved_files = files - if files and not agentic: - resolved_files = {} - for label, f in files.items(): + if resolved_files and not agentic: + env_resolved_files = {} + for label, f in resolved_files.items(): if f.source == "env" and f._env is not None: b64 = await _collect_file_from_env_async(f._env, f.filename) if b64 is not None: - resolved_files[label] = File.from_base64( + env_resolved_files[label] = File.from_base64( b64, f.filename or "file", _guess_file_media_type(f.filename or "file"), ) else: # Async collection failed — use collect source directly - resolved_files[label] = File( + env_resolved_files[label] = File( source="collect", filename=f.filename, ) else: - resolved_files[label] = f + env_resolved_files[label] = f + resolved_files = env_resolved_files + _print_judge_call_start(rubric, resolved_images, agentic, model, files=resolved_files) + + if self._llm_provider is not None: + # Route through pluggable LLM provider + from ..llm_provider import GradeRequest + + request = GradeRequest( + rubric=rubric, + submission=submission, + ground_truth=ground_truth, + problem=problem, + context=effective_context, + conversation=conversation, + images=resolved_images, + files=resolved_files, + model=model, + provider=provider, + agentic=agentic, + collect=collect, + task_id=task_id, + instance_id=self._instance_id, + ) + grade_response = await self._llm_provider.agrade(request) + return _parse_grade_response(grade_response.to_dict()) + + # Default: route through Fleet orchestrator body = _build_grade_request( self._instance_id, rubric, submission, ground_truth=ground_truth, problem=problem, - context=context, - reference_claims=reference_claims, + context=effective_context, + reference_claims=None, # already folded into context conversation=conversation, images=resolved_images, files=resolved_files, @@ -147,6 +218,5 @@ async def grade( task_id=task_id, ) - _print_judge_call_start(rubric, resolved_images, agentic, model, files=resolved_files) response = await self._client.request("POST", "/v1/judge/grade", json=body) return _parse_grade_response(response.json()) diff --git a/fleet/client.py b/fleet/client.py index d01ee4ea..b54d3550 100644 --- a/fleet/client.py +++ b/fleet/client.py @@ -60,6 +60,7 @@ if TYPE_CHECKING: from .verifiers import SyncVerifierFunction from .judge import SyncJudge + from .llm_provider import LLMProvider def _json_default(x: Any) -> Any: @@ -344,12 +345,19 @@ def message_count(self) -> int: class SyncEnv(EnvironmentBase): - def __init__(self, client: Optional[SyncWrapper], **kwargs): + def __init__( + self, + client: Optional[SyncWrapper], + *, + llm_provider: Optional["LLMProvider"] = None, + **kwargs, + ): super().__init__(**kwargs) self._client = client self._apps: Dict[str, InstanceClient] = {} self._instance: Optional[InstanceClient] = None self._judge: Optional["SyncJudge"] = None + self._llm_provider = llm_provider self._manager_url_override: Optional[str] = None # For URL mode @property @@ -435,13 +443,18 @@ def mcp(self) -> SyncMCPResource: @property def judge(self) -> "SyncJudge": - """LLM-as-judge grading via orchestrator API.""" + """LLM-as-judge grading. + + Routes through Fleet orchestrator by default. Set ``llm_provider`` + on the environment to route to an external provider instead. + """ if self._judge is None: from .judge import SyncJudge self._judge = SyncJudge( - client=self._load_client, + client=self._client, instance_id=self.instance_id, + llm_provider=self._llm_provider, ) return self._judge diff --git a/fleet/judge.py b/fleet/judge.py index c660c6d6..03ec2b36 100644 --- a/fleet/judge.py +++ b/fleet/judge.py @@ -1,10 +1,21 @@ -"""Fleet SDK Judge - LLM-as-Judge grading via orchestrator API. +"""Fleet SDK Judge - LLM-as-Judge grading. Provides env.judge.grade() for verifier scripts to grade submissions using LLM judges without managing API keys, HTTP calls, or response parsing. -All LLM calls happen server-side on the orchestrator — the SDK just sends -the rubric, submission, and artifacts, and gets back a score. +By default, LLM calls route through the Fleet orchestrator. For on-prem +or external deployments, pass an ``llm_provider`` to route calls to +any OpenAI-compatible endpoint (OpenRouter, Anthropic, local models, etc.):: + + from fleet.llm_provider import ExternalProvider + + provider = ExternalProvider( + api_key="sk-or-...", + base_url="https://openrouter.ai/api/v1", + model="anthropic/claude-sonnet-4", + ) + judge = SyncJudge(client=None, instance_id="local", llm_provider=provider) + result = judge.grade(rubric, submission) """ import base64 @@ -16,6 +27,7 @@ if TYPE_CHECKING: from .base import SyncWrapper + from .llm_provider import LLMProvider logger = logging.getLogger(__name__) @@ -143,11 +155,21 @@ def serialize(self) -> dict: class Image: """Reference to an image for LLM judge grading. - Use the static constructors to create instances: + Preferred constructor (source-agnostic):: + + Image.from_path("screenshots/gold.png") + Image.from_path("s3://bucket/key.png") + Image.from_path("https://example.com/img.png") + + The LLM provider resolves the path at grade-time, so verifier code + stays independent of storage backends. + + Legacy constructors (still supported for backward compat): Image.s3("s3://bucket/key") - S3 URL, fetched server-side Image.from_url("https://...") - HTTP URL, fetched server-side Image.from_base64(data, "file.png") - Inline base64 data Image.from_env(env, "plot.png") - Collect from environment + Image.from_local("/path/to/image.png") - Read from local filesystem """ def __init__( @@ -159,6 +181,8 @@ def __init__( filename: Optional[str] = None, media_type: Optional[str] = None, _env: Optional[Any] = None, + _local_path: Optional[str] = None, + _path: Optional[str] = None, ): self.source = source self.url = url @@ -166,6 +190,8 @@ def __init__( self.filename = filename self.media_type = media_type self._env = _env + self._local_path = _local_path + self._path = _path @staticmethod def s3(url: str, media_type: Optional[str] = None) -> "Image": @@ -200,6 +226,66 @@ def from_env(env: Any, filename: str) -> "Image": """ return Image(source="env", filename=filename, _env=env) + @staticmethod + def from_local(path: str, media_type: Optional[str] = None) -> "Image": + """Read an image from a local file path. + + The file is read and base64-encoded at serialization time (lazy), so + the path must be valid when ``serialize()`` or an external provider + processes the image. For on-prem deployments this is the primary way + to supply reference images without S3 access. + + Args: + path: Absolute or relative path to the image file. + media_type: Optional MIME type override (auto-detected from extension). + """ + fname = os.path.basename(path) + return Image( + source="local", + filename=fname, + media_type=media_type or _guess_media_type(fname), + _local_path=path, + ) + + @staticmethod + def from_path(path: str, media_type: Optional[str] = None) -> "Image": + """Create a source-agnostic image reference from a path or URI. + + The path is resolved at grade-time by the LLM provider's + ``resolve_image()`` method. This keeps verifier code independent + of storage backends (S3, local filesystem, HTTP, etc.). + + Examples:: + + Image.from_path("screenshots/gold.png") # relative + Image.from_path("/abs/path/to/gold.png") # absolute local + Image.from_path("s3://bucket/screenshots/gold.png") # S3 URI + Image.from_path("https://example.com/gold.png") # HTTP URL + + Args: + path: Path or URI to the image. Scheme detection and fetching + are handled by the provider at resolve-time. + media_type: Optional MIME type override (auto-detected from extension). + """ + fname = os.path.basename(path) + return Image( + source="path", + filename=fname, + media_type=media_type or _guess_media_type(fname), + _path=path, + ) + + def _resolve_local(self) -> Optional[str]: + """Read the local file and return base64 data, or None on failure.""" + if not self._local_path: + return None + try: + with open(self._local_path, "rb") as f: + return base64.b64encode(f.read()).decode("ascii") + except (OSError, IOError) as e: + logger.warning("Failed to read local image %s: %s", self._local_path, e) + return None + def serialize(self, *, label: Optional[str] = None, agentic: bool = False) -> dict: """Serialize for the orchestrator API request body.""" d: dict @@ -217,6 +303,39 @@ def serialize(self, *, label: Optional[str] = None, agentic: bool = False) -> di "data": self.data, "media_type": self.media_type or _guess_media_type(self.filename or "image.png"), } + elif self.source == "local": + b64 = self._resolve_local() + if b64 is not None: + d = { + "source": "base64", + "data": b64, + "media_type": self.media_type or _guess_media_type(self.filename or "image.png"), + } + else: + raise ValueError(f"Cannot read local image: {self._local_path}") + elif self.source == "path": + # Auto-detect scheme as fallback when no provider resolved this. + path = self._path or self.filename or "" + if path.startswith("s3://"): + d = {"source": "s3", "url": path} + if self.media_type: + d["media_type"] = self.media_type + elif path.startswith(("http://", "https://")): + d = {"source": "url", "url": path} + if self.media_type: + d["media_type"] = self.media_type + else: + # Treat as local file + self._local_path = path + b64 = self._resolve_local() + if b64 is not None: + d = { + "source": "base64", + "data": b64, + "media_type": self.media_type or _guess_media_type(self.filename or "image.png"), + } + else: + raise ValueError(f"Cannot read image path: {path}") elif self.source == "collect": d = {"source": "collect", "selector": self.filename} elif self.source == "env": @@ -243,11 +362,19 @@ def serialize(self, *, label: Optional[str] = None, agentic: bool = False) -> di class File: """Reference to an arbitrary file for LLM judge grading. - Supports any file type (PDF, CSV, STEP, STL, etc.) via the Anthropic - Files API. Use the static constructors to create instances: + Preferred constructor (source-agnostic):: + + File.from_path("reports/output.pdf") + File.from_path("s3://bucket/data.csv") + + The LLM provider resolves the path at grade-time, so verifier code + stays independent of storage backends. + + Legacy constructors (still supported for backward compat): File.s3("s3://bucket/key") - S3 URL, fetched server-side File.from_base64(data, "part.step", "application/step") - Inline base64 data File.from_env(env, "exported_part.step") - Collect from environment + File.from_local("/path/to/file.pdf") - Read from local filesystem """ def __init__( @@ -259,6 +386,8 @@ def __init__( filename: Optional[str] = None, media_type: Optional[str] = None, _env: Optional[Any] = None, + _local_path: Optional[str] = None, + _path: Optional[str] = None, ): self.source = source self.url = url @@ -266,6 +395,8 @@ def __init__( self.filename = filename self.media_type = media_type self._env = _env + self._local_path = _local_path + self._path = _path @staticmethod def s3(url: str, media_type: Optional[str] = None) -> "File": @@ -295,6 +426,65 @@ def from_env(env: Any, filename: str) -> "File": """ return File(source="env", filename=filename, _env=env) + @staticmethod + def from_local(path: str, media_type: Optional[str] = None) -> "File": + """Read a file from a local file path. + + The file is read and base64-encoded at serialization time (lazy), so + the path must be valid when ``serialize()`` or an external provider + processes the file. For on-prem deployments this is the primary way + to supply reference files without S3 access. + + Args: + path: Absolute or relative path to the file. + media_type: Optional MIME type override (auto-detected from extension). + """ + fname = os.path.basename(path) + return File( + source="local", + filename=fname, + media_type=media_type or _guess_file_media_type(fname), + _local_path=path, + ) + + @staticmethod + def from_path(path: str, media_type: Optional[str] = None) -> "File": + """Create a source-agnostic file reference from a path or URI. + + The path is resolved at grade-time by the LLM provider's + ``resolve_file()`` method. This keeps verifier code independent + of storage backends (S3, local filesystem, HTTP, etc.). + + Examples:: + + File.from_path("reports/output.pdf") # relative + File.from_path("/abs/path/to/data.csv") # absolute local + File.from_path("s3://bucket/reports/output.pdf") # S3 URI + + Args: + path: Path or URI to the file. Scheme detection and fetching + are handled by the provider at resolve-time. + media_type: Optional MIME type override (auto-detected from extension). + """ + fname = os.path.basename(path) + return File( + source="path", + filename=fname, + media_type=media_type or _guess_file_media_type(fname), + _path=path, + ) + + def _resolve_local(self) -> Optional[str]: + """Read the local file and return base64 data, or None on failure.""" + if not self._local_path: + return None + try: + with open(self._local_path, "rb") as f: + return base64.b64encode(f.read()).decode("ascii") + except (OSError, IOError) as e: + logger.warning("Failed to read local file %s: %s", self._local_path, e) + return None + def serialize(self, *, label: Optional[str] = None, agentic: bool = False) -> dict: """Serialize for the orchestrator API request body.""" d: dict @@ -309,6 +499,37 @@ def serialize(self, *, label: Optional[str] = None, agentic: bool = False) -> di "filename": self.filename, "media_type": self.media_type or _guess_file_media_type(self.filename or "file"), } + elif self.source == "local": + b64 = self._resolve_local() + if b64 is not None: + d = { + "source": "base64", + "data": b64, + "filename": self.filename, + "media_type": self.media_type or _guess_file_media_type(self.filename or "file"), + } + else: + raise ValueError(f"Cannot read local file: {self._local_path}") + elif self.source == "path": + # Auto-detect scheme as fallback when no provider resolved this. + path = self._path or self.filename or "" + if path.startswith("s3://"): + d = {"source": "s3", "url": path} + if self.media_type: + d["media_type"] = self.media_type + else: + # Treat as local file + self._local_path = path + b64 = self._resolve_local() + if b64 is not None: + d = { + "source": "base64", + "data": b64, + "filename": self.filename, + "media_type": self.media_type or _guess_file_media_type(self.filename or "file"), + } + else: + raise ValueError(f"Cannot read file path: {path}") elif self.source == "collect": d = {"source": "collect", "selector": self.filename} elif self.source == "env": @@ -942,16 +1163,55 @@ def _print_judge_result(data: dict) -> None: # --------------------------------------------------------------------------- +_UNSET = object() # sentinel — distinguishes "not passed" from None + + class SyncJudge: - """LLM-as-judge grading — calls orchestrator API, not environment API. + """LLM-as-judge grading. + + Accessed as ``env.judge`` on SyncEnv instances. + + Provider resolution order: + + 1. Explicit ``llm_provider`` kwarg (highest priority) + 2. ``FLEET_LLM_API_KEY`` env var → auto-builds ``ExternalProvider`` + 3. Fleet orchestrator (default fallback) - Accessed as env.judge on SyncEnv instances. + Set env vars for zero-code external routing:: + + export FLEET_LLM_API_KEY="sk-or-..." + export FLEET_LLM_BASE_URL="https://openrouter.ai/api/v1" + export FLEET_LLM_MODEL="anthropic/claude-sonnet-4" + + Or pass a provider explicitly:: + + from fleet.llm_provider import ExternalProvider + + provider = ExternalProvider( + api_key="sk-or-...", + base_url="https://openrouter.ai/api/v1", + model="anthropic/claude-sonnet-4", + ) + judge = SyncJudge(client=None, instance_id="local", llm_provider=provider) """ - def __init__(self, client: "SyncWrapper", instance_id: str): + def __init__( + self, + client: Optional["SyncWrapper"], + instance_id: str, + *, + llm_provider: Any = _UNSET, + ): self._client = client self._instance_id = instance_id + if llm_provider is _UNSET: + # Auto-detect from env vars; None means "use Fleet orchestrator" + from .llm_provider import resolve_provider + self._llm_provider = resolve_provider() + else: + self._llm_provider = llm_provider + def grade( self, rubric: Union[str, Rubric], @@ -970,7 +1230,11 @@ def grade( collect: Optional[Dict[str, List[str]]] = None, task_id: Optional[str] = None, ) -> JudgeResult: - """Grade a submission using LLM-as-judge via the orchestrator API. + """Grade a submission using LLM-as-judge. + + Routes through the Fleet orchestrator by default. If an + ``llm_provider`` was set at construction time, calls the external + provider directly instead. Returns a JudgeResult (float subclass with .details, .criteria, .feedback) that can be returned directly from a verifier function. @@ -991,17 +1255,59 @@ def grade( collect: File patterns for orchestrator to collect (agentic mode). task_id: Optional task ID for tracking. """ + # Fold reference_claims into context (shared logic regardless of provider) + effective_context = context + if reference_claims is not None: + if effective_context: + effective_context = f"{effective_context}\n\n## Reference Claims\n{reference_claims}" + else: + effective_context = f"## Reference Claims\n{reference_claims}" + + # Resolve path-based images/files through the provider + resolved_images = images + resolved_files = files + if self._llm_provider is not None: + resolved_images = self._llm_provider.resolve_images(images) + resolved_files = self._llm_provider.resolve_files(files) + + _print_judge_call_start(rubric, resolved_images, agentic, model, files=resolved_files) + + if self._llm_provider is not None: + # Route through pluggable LLM provider + from .llm_provider import GradeRequest + + request = GradeRequest( + rubric=rubric, + submission=submission, + ground_truth=ground_truth, + problem=problem, + context=effective_context, + conversation=conversation, + images=resolved_images, + files=resolved_files, + model=model, + provider=provider, + agentic=agentic, + collect=collect, + task_id=task_id, + instance_id=self._instance_id, + ) + grade_response = self._llm_provider.grade(request) + return _parse_grade_response(grade_response.to_dict()) + + # Default: route through Fleet orchestrator + # (path-based images/files auto-resolve in serialize() fallback) body = _build_grade_request( self._instance_id, rubric, submission, ground_truth=ground_truth, problem=problem, - context=context, - reference_claims=reference_claims, + context=effective_context, + reference_claims=None, # already folded into context conversation=conversation, - images=images, - files=files, + images=resolved_images, + files=resolved_files, model=model, provider=provider, agentic=agentic, @@ -1009,6 +1315,5 @@ def grade( task_id=task_id, ) - _print_judge_call_start(rubric, images, agentic, model, files=files) response = self._client.request("POST", "/v1/judge/grade", json=body) return _parse_grade_response(response.json()) diff --git a/fleet/llm_provider.py b/fleet/llm_provider.py new file mode 100644 index 00000000..6c8acbc1 --- /dev/null +++ b/fleet/llm_provider.py @@ -0,0 +1,876 @@ +"""LLM Provider abstraction for Fleet SDK Judge. + +Defines a pluggable interface for routing LLM judge calls to either +the Fleet orchestrator (default) or external providers like OpenRouter, +Anthropic API, etc. This enables on-prem deployments that don't depend +on Fleet's internal orchestrator endpoints. + +Configuration via environment variables (auto-detected at judge init):: + + # Set these env vars to route judge calls to an external provider. + # When FLEET_LLM_API_KEY is unset, calls route through Fleet orchestrator. + export FLEET_LLM_API_KEY="sk-or-..." # required + export FLEET_LLM_BASE_URL="https://openrouter.ai/api/v1" # optional (default: OpenRouter) + export FLEET_LLM_MODEL="anthropic/claude-sonnet-4" # optional (default: anthropic/claude-sonnet-4) + export FLEET_LLM_TEMPERATURE="0.0" # optional (default: 0.0) + export FLEET_LLM_MAX_TOKENS="4096" # optional (default: 4096) + export FLEET_LLM_TIMEOUT="300" # optional (default: 300s) + +Or configure programmatically:: + + from fleet.llm_provider import ExternalProvider + + provider = ExternalProvider( + api_key="sk-or-...", + base_url="https://openrouter.ai/api/v1", + model="anthropic/claude-sonnet-4", + ) + judge = SyncJudge(client=None, instance_id="local", llm_provider=provider) + result = judge.grade(rubric, submission) +""" + +import json +import logging +import os +import time +from abc import ABC, abstractmethod +from dataclasses import dataclass, field +from typing import Any, Dict, List, Optional, Union, TYPE_CHECKING + +import httpx + +if TYPE_CHECKING: + from .judge import Rubric, Image, File + +logger = logging.getLogger(__name__) + + +# --------------------------------------------------------------------------- +# Environment variable names +# --------------------------------------------------------------------------- + +ENV_LLM_API_KEY = "FLEET_LLM_API_KEY" +ENV_LLM_BASE_URL = "FLEET_LLM_BASE_URL" +ENV_LLM_MODEL = "FLEET_LLM_MODEL" +ENV_LLM_TEMPERATURE = "FLEET_LLM_TEMPERATURE" +ENV_LLM_MAX_TOKENS = "FLEET_LLM_MAX_TOKENS" +ENV_LLM_TIMEOUT = "FLEET_LLM_TIMEOUT" + + +# --------------------------------------------------------------------------- +# Grade request / response types (provider-agnostic) +# --------------------------------------------------------------------------- + + +@dataclass +class GradeRequest: + """Provider-agnostic grading request. + + Contains all the information needed to perform LLM-as-judge grading, + independent of whether the call routes through Fleet or an external API. + """ + + rubric: Union[str, "Rubric"] + submission: Optional[str] = None + ground_truth: Optional[Union[str, dict]] = None + problem: Optional[str] = None + context: Optional[str] = None + conversation: Optional[List[dict]] = None + images: Optional[Dict[str, "Image"]] = None + files: Optional[Dict[str, "File"]] = None + model: Optional[str] = None + provider: Optional[str] = None + agentic: bool = False + collect: Optional[Dict[str, List[str]]] = None + task_id: Optional[str] = None + instance_id: Optional[str] = None + + +@dataclass +class GradeResponse: + """Provider-agnostic grading response. + + Normalized structure returned by all providers. Maps to the existing + JudgeResult construction in ``_parse_grade_response``. + """ + + normalized_score: float + total_score: float = 0.0 + max_score: float = 0.0 + criteria: List[dict] = field(default_factory=list) + feedback: str = "" + model_used: str = "" + provider_used: str = "" + accumulators: Optional[dict] = None + raw: Optional[dict] = None # Full raw response for pass-through + + def to_dict(self) -> dict: + """Convert to dict matching the orchestrator response schema.""" + d: dict = { + "normalized_score": self.normalized_score, + "total_score": self.total_score, + "max_score": self.max_score, + "model_used": self.model_used, + "provider_used": self.provider_used, + } + if self.criteria: + d["criteria"] = self.criteria + if self.feedback: + d["feedback"] = self.feedback + if self.accumulators: + d["accumulators"] = self.accumulators + return d + + +# --------------------------------------------------------------------------- +# Abstract base +# --------------------------------------------------------------------------- + + +class LLMProvider(ABC): + """Abstract interface for LLM judge backends. + + Implementations must provide ``grade()`` (sync) and/or ``agrade()`` + (async). The judge classes call whichever variant matches their + execution model. + + Providers can override ``resolve_image()`` / ``resolve_file()`` to + customise how ``source="path"`` references are resolved (e.g. prepend + an S3 prefix, fetch from GCS, etc.). The default implementation + auto-detects the URI scheme and delegates to the appropriate legacy + constructor. + """ + + # ------------------------------------------------------------------ + # Path resolution (override for custom storage backends) + # ------------------------------------------------------------------ + + def resolve_image(self, image: Any) -> Any: + """Resolve a path-based Image to a concrete source. + + Override in subclasses for custom resolution logic (e.g. prepending + an S3 bucket prefix, fetching from GCS/Azure Blob, etc.). + + The default implementation auto-detects the URI scheme: + + - ``s3://`` → ``Image.s3()`` + - ``http(s)://`` → ``Image.from_url()`` + - anything else → ``Image.from_local()`` + + Non-path images are returned unchanged. + """ + if getattr(image, "source", None) != "path": + return image + + from .judge import Image as _Image + + path = image._path or image.filename or "" + mt = image.media_type + + if path.startswith("s3://"): + return _Image.s3(path, media_type=mt) + elif path.startswith(("http://", "https://")): + return _Image.from_url(path, media_type=mt) + else: + return _Image.from_local(path, media_type=mt) + + def resolve_file(self, file: Any) -> Any: + """Resolve a path-based File to a concrete source. + + Same semantics as ``resolve_image()`` but for File objects. + """ + if getattr(file, "source", None) != "path": + return file + + from .judge import File as _File + + path = file._path or file.filename or "" + mt = file.media_type + + if path.startswith("s3://"): + return _File.s3(path, media_type=mt) + else: + return _File.from_local(path, media_type=mt) + + def resolve_images(self, images: Optional[Dict[str, Any]]) -> Optional[Dict[str, Any]]: + """Resolve all path-based images in a dict.""" + if not images: + return images + return {label: self.resolve_image(img) for label, img in images.items()} + + def resolve_files(self, files: Optional[Dict[str, Any]]) -> Optional[Dict[str, Any]]: + """Resolve all path-based files in a dict.""" + if not files: + return files + return {label: self.resolve_file(f) for label, f in files.items()} + + # ------------------------------------------------------------------ + # Grading (must implement) + # ------------------------------------------------------------------ + + @abstractmethod + def grade(self, request: GradeRequest) -> GradeResponse: + """Execute a synchronous grading call.""" + ... + + async def agrade(self, request: GradeRequest) -> GradeResponse: + """Execute an asynchronous grading call. + + Default implementation runs the sync ``grade()`` method in a thread + executor so it does not block the event loop. Providers that support + native async should override this. + """ + import asyncio + + loop = asyncio.get_running_loop() + return await loop.run_in_executor(None, self.grade, request) + + +# --------------------------------------------------------------------------- +# Fleet orchestrator provider (default) +# --------------------------------------------------------------------------- + + +class FleetProvider(LLMProvider): + """Routes judge calls through the Fleet orchestrator API. + + This is the default provider — it preserves the existing behavior where + ``SyncJudge.grade()`` calls ``POST /v1/judge/grade`` on the orchestrator. + """ + + def __init__(self, client: Any, instance_id: str): + self._client = client + self._instance_id = instance_id + + def grade(self, request: GradeRequest) -> GradeResponse: + from .judge import _build_grade_request + + body = _build_grade_request( + self._instance_id, + request.rubric, + request.submission, + ground_truth=request.ground_truth, + problem=request.problem, + context=request.context, + conversation=request.conversation, + images=request.images, + files=request.files, + model=request.model, + provider=request.provider, + agentic=request.agentic, + collect=request.collect, + task_id=request.task_id, + ) + + response = self._client.request("POST", "/v1/judge/grade", json=body) + data = response.json() + + return GradeResponse( + normalized_score=float(data.get("normalized_score", 0.0)), + total_score=float(data.get("total_score", 0)), + max_score=float(data.get("max_score", 0)), + criteria=data.get("criteria", []), + feedback=data.get("feedback", ""), + model_used=data.get("model_used", ""), + provider_used=data.get("provider_used", ""), + accumulators=data.get("accumulators"), + raw=data, + ) + + async def agrade(self, request: GradeRequest) -> GradeResponse: + from .judge import _build_grade_request + + body = _build_grade_request( + self._instance_id, + request.rubric, + request.submission, + ground_truth=request.ground_truth, + problem=request.problem, + context=request.context, + conversation=request.conversation, + images=request.images, + files=request.files, + model=request.model, + provider=request.provider, + agentic=request.agentic, + collect=request.collect, + task_id=request.task_id, + ) + + response = await self._client.request("POST", "/v1/judge/grade", json=body) + data = response.json() + + return GradeResponse( + normalized_score=float(data.get("normalized_score", 0.0)), + total_score=float(data.get("total_score", 0)), + max_score=float(data.get("max_score", 0)), + criteria=data.get("criteria", []), + feedback=data.get("feedback", ""), + model_used=data.get("model_used", ""), + provider_used=data.get("provider_used", ""), + accumulators=data.get("accumulators"), + raw=data, + ) + + +# --------------------------------------------------------------------------- +# External provider (OpenRouter, Anthropic, etc.) +# --------------------------------------------------------------------------- + +# Default system prompt for the judge when running externally +_DEFAULT_JUDGE_SYSTEM_PROMPT = """\ +You are an expert judge evaluating a submission against a rubric. +You must evaluate the submission fairly and provide a score for each criterion. + +Respond with a JSON object in this exact format: +{ + "criteria": [ + { + "name": "", + "score": , + "max_score": , + "reasoning": "" + } + ], + "feedback": "" +} + +IMPORTANT: Respond ONLY with the JSON object. No markdown fences, no extra text.""" + + +def _build_judge_user_message(request: GradeRequest) -> str: + """Build the user message content for the judge LLM call.""" + from .judge import Rubric + + parts: List[str] = [] + + # Problem statement + if request.problem: + parts.append(f"## Problem\n{request.problem}") + + # Rubric + if isinstance(request.rubric, str): + parts.append(f"## Rubric\n{request.rubric}") + elif isinstance(request.rubric, Rubric): + rubric_lines = [] + for c in request.rubric.criteria: + rubric_lines.append(f"- **{c.name}** (max {c.max} points): {c._render_description()}") + parts.append(f"## Rubric\n" + "\n".join(rubric_lines)) + + # Ground truth + if request.ground_truth: + gt = request.ground_truth + if isinstance(gt, dict): + gt = json.dumps(gt, indent=2) + parts.append(f"## Ground Truth / Expected Answer\n{gt}") + + # Context + if request.context: + parts.append(f"## Additional Context\n{request.context}") + + # Conversation history + if request.conversation: + conv_lines = [] + for msg in request.conversation: + role = msg.get("role", "unknown") + content = msg.get("content", "") + conv_lines.append(f"[{role}]: {content}") + parts.append(f"## Conversation History\n" + "\n\n".join(conv_lines)) + + # Submission + if request.submission: + parts.append(f"## Submission to Grade\n{request.submission}") + else: + parts.append("## Submission to Grade\n(No submission text provided)") + + return "\n\n".join(parts) + + +def _build_anthropic_messages( + request: GradeRequest, + system_prompt: str, +) -> tuple: + """Build messages list for the Anthropic/OpenAI chat completions format. + + Returns (system_prompt, messages) tuple. + """ + from .judge import Rubric + + # Use rubric's system_prompt override if provided + if isinstance(request.rubric, Rubric) and request.rubric.system_prompt: + system_prompt = request.rubric.system_prompt + + user_content: list = [] + + # Add images as base64 content blocks (Anthropic vision format) + if request.images: + for label, img in request.images.items(): + # Resolve path-based images to base64 (fallback if not pre-resolved) + if img.source == "path" and img._path: + path = img._path + if path.startswith(("http://", "https://")): + user_content.append({ + "type": "image", + "source": { + "type": "url", + "url": path, + }, + }) + else: + # Treat as local file + import base64 as _b64 + try: + with open(path, "rb") as fh: + b64 = _b64.b64encode(fh.read()).decode("ascii") + user_content.append({ + "type": "image", + "source": { + "type": "base64", + "media_type": img.media_type or "image/png", + "data": b64, + }, + }) + except (OSError, IOError): + logger.warning("Skipping unreadable path image: %s", path) + continue + + # Resolve local images to base64 first + if img.source == "local" and img._local_path: + b64 = img._resolve_local() + if b64: + user_content.append({ + "type": "image", + "source": { + "type": "base64", + "media_type": img.media_type or "image/png", + "data": b64, + }, + }) + else: + logger.warning("Skipping unreadable local image: %s", img._local_path) + continue + + if img.data: # base64 data available + user_content.append({ + "type": "image", + "source": { + "type": "base64", + "media_type": img.media_type or "image/png", + "data": img.data, + }, + }) + elif img.url: + user_content.append({ + "type": "image", + "source": { + "type": "url", + "url": img.url, + }, + }) + + # Add file contents as text blocks + if request.files: + _TEXT_MEDIA_TYPES = { + "text/plain", "text/markdown", "text/html", "text/csv", + "text/tab-separated-values", "application/json", "application/xml", + "application/x-yaml", + } + for label, f in request.files.items(): + file_text = _resolve_file_text(f) + if file_text is not None: + mt = getattr(f, "media_type", "") or "" + if mt in _TEXT_MEDIA_TYPES or _is_likely_text(file_text): + user_content.append({ + "type": "text", + "text": f"## File: {label} ({getattr(f, 'filename', label)})\n\n{file_text}", + }) + else: + # Binary file — include metadata only + user_content.append({ + "type": "text", + "text": ( + f"## File: {label} ({getattr(f, 'filename', label)})\n\n" + f"[Binary file, media_type={mt or 'application/octet-stream'}]" + ), + }) + else: + logger.warning("Skipping unresolvable file: %s", label) + + # Add text content + user_text = _build_judge_user_message(request) + user_content.append({"type": "text", "text": user_text}) + + messages = [{"role": "user", "content": user_content}] + return system_prompt, messages + + +def _resolve_file_text(f: Any) -> Optional[str]: + """Try to extract text content from a File object. + + Returns the decoded text, or None if the file cannot be resolved. + """ + import base64 as _b64 + + # Already has base64 data + if getattr(f, "data", None): + try: + return _b64.b64decode(f.data).decode("utf-8") + except (UnicodeDecodeError, Exception): + return _b64.b64decode(f.data).decode("latin-1", errors="replace") + + # Local file + if getattr(f, "source", None) == "local" and getattr(f, "_local_path", None): + try: + with open(f._local_path, "rb") as fh: + raw = fh.read() + try: + return raw.decode("utf-8") + except UnicodeDecodeError: + return raw.decode("latin-1", errors="replace") + except (OSError, IOError): + return None + + # Path-based file — read from disk + if getattr(f, "source", None) == "path" and getattr(f, "_path", None): + path = f._path + if path.startswith("s3://") or path.startswith(("http://", "https://")): + # Can't read remote files inline — skip + return None + try: + with open(path, "rb") as fh: + raw = fh.read() + try: + return raw.decode("utf-8") + except UnicodeDecodeError: + return raw.decode("latin-1", errors="replace") + except (OSError, IOError): + return None + + return None + + +def _is_likely_text(content: str) -> bool: + """Heuristic: check if content looks like readable text.""" + if not content: + return False + # If >95% of first 1024 chars are printable ASCII/unicode, it's text + sample = content[:1024] + printable = sum(1 for c in sample if c.isprintable() or c in "\n\r\t") + return printable / len(sample) > 0.95 + + +def _parse_llm_judge_response( + raw_text: str, + rubric: Any, + model: str, + provider: str, +) -> GradeResponse: + """Parse the LLM's JSON response into a GradeResponse.""" + from .judge import Rubric + + # Strip markdown fences if present + text = raw_text.strip() + if text.startswith("```"): + # Remove opening fence (```json or ```) + text = text.split("\n", 1)[1] if "\n" in text else text[3:] + if text.endswith("```"): + text = text[:-3] + text = text.strip() + + try: + data = json.loads(text) + except json.JSONDecodeError as e: + logger.warning("Failed to parse judge LLM response as JSON: %s", e) + logger.debug("Raw response: %s", raw_text[:500]) + return GradeResponse( + normalized_score=0.0, + feedback=f"Failed to parse judge response: {e}", + model_used=model, + provider_used=provider, + ) + + criteria = data.get("criteria", []) + feedback = data.get("feedback", "") + + # Compute scores + total_score = sum(c.get("score", 0) for c in criteria) + if isinstance(rubric, Rubric): + max_score = rubric.max_score + else: + max_score = sum(c.get("max_score", 0) for c in criteria) + + normalized = total_score / max_score if max_score > 0 else 0.0 + + return GradeResponse( + normalized_score=normalized, + total_score=float(total_score), + max_score=float(max_score), + criteria=criteria, + feedback=feedback, + model_used=model, + provider_used=provider, + raw=data, + ) + + +class ExternalProvider(LLMProvider): + """Routes judge calls directly to an external LLM API. + + Supports any OpenAI-compatible chat completions endpoint, including: + - OpenRouter (https://openrouter.ai/api/v1) + - Anthropic via proxy + - Azure OpenAI + - Local models (vLLM, Ollama, etc.) + + Args: + api_key: API key for the provider. + base_url: Base URL for the chat completions API. + Defaults to OpenRouter. + model: Model identifier (e.g., "anthropic/claude-sonnet-4"). + system_prompt: Override the default judge system prompt. + timeout: Request timeout in seconds (default: 300). + extra_headers: Additional headers to include in requests. + temperature: Sampling temperature (default: 0.0 for deterministic). + max_tokens: Maximum tokens in response (default: 4096). + + Example:: + + provider = ExternalProvider( + api_key="sk-or-...", + base_url="https://openrouter.ai/api/v1", + model="anthropic/claude-sonnet-4", + ) + judge = SyncJudge(client=None, instance_id="local", llm_provider=provider) + result = judge.grade(rubric, submission) + """ + + DEFAULT_BASE_URL = "https://openrouter.ai/api/v1" + DEFAULT_MODEL = "anthropic/claude-sonnet-4" + + def __init__( + self, + *, + api_key: str, + base_url: Optional[str] = None, + model: Optional[str] = None, + system_prompt: Optional[str] = None, + timeout: float = 300.0, + extra_headers: Optional[Dict[str, str]] = None, + temperature: float = 0.0, + max_tokens: int = 4096, + ): + if not api_key or not api_key.strip(): + raise ValueError("ExternalProvider requires a non-empty api_key") + self.api_key = api_key + self.base_url = (base_url or self.DEFAULT_BASE_URL).rstrip("/") + self.model = model or self.DEFAULT_MODEL + self.system_prompt = system_prompt or _DEFAULT_JUDGE_SYSTEM_PROMPT + self.timeout = timeout + self.extra_headers = extra_headers or {} + self.temperature = temperature + self.max_tokens = max_tokens + + def _get_headers(self) -> Dict[str, str]: + headers = { + "Authorization": f"Bearer {self.api_key}", + "Content-Type": "application/json", + } + headers.update(self.extra_headers) + return headers + + def _build_request_body(self, request: GradeRequest) -> dict: + """Build an OpenAI-compatible chat completions request body.""" + model = request.model or self.model + system_prompt, messages = _build_anthropic_messages( + request, self.system_prompt, + ) + + # Convert to OpenAI chat format + oai_messages: List[dict] = [ + {"role": "system", "content": system_prompt}, + ] + for msg in messages: + content = msg["content"] + if isinstance(content, list): + # Convert multimodal content blocks to OpenAI format + oai_content: List[dict] = [] + for block in content: + if block.get("type") == "text": + oai_content.append({"type": "text", "text": block["text"]}) + elif block.get("type") == "image": + source = block.get("source", {}) + if source.get("type") == "base64": + oai_content.append({ + "type": "image_url", + "image_url": { + "url": f"data:{source.get('media_type', 'image/png')};base64,{source['data']}", + }, + }) + elif source.get("type") == "url": + oai_content.append({ + "type": "image_url", + "image_url": {"url": source["url"]}, + }) + oai_messages.append({"role": msg["role"], "content": oai_content}) + else: + oai_messages.append(msg) + + return { + "model": model, + "messages": oai_messages, + "temperature": self.temperature, + "max_tokens": self.max_tokens, + } + + def grade(self, request: GradeRequest) -> GradeResponse: + """Grade via external OpenAI-compatible API (sync).""" + model = request.model or self.model + body = self._build_request_body(request) + url = f"{self.base_url}/chat/completions" + + start_ms = time.time() * 1000 + + try: + with httpx.Client(timeout=self.timeout) as client: + response = client.post(url, json=body, headers=self._get_headers()) + response.raise_for_status() + except httpx.HTTPStatusError as e: + logger.warning("LLM API HTTP error: %s %s", e.response.status_code, e.response.text[:500]) + return GradeResponse( + normalized_score=0.0, + feedback=f"LLM API error: {e.response.status_code} - {e.response.text[:500]}", + model_used=model, + provider_used=self.base_url, + ) + except httpx.RequestError as e: + logger.warning("LLM API request failed: %s", e) + return GradeResponse( + normalized_score=0.0, + feedback=f"LLM API request failed: {e}", + model_used=model, + provider_used=self.base_url, + ) + + elapsed_ms = time.time() * 1000 - start_ms + data = response.json() + + # Extract response text from OpenAI format + raw_text = data["choices"][0]["message"]["content"] + + result = _parse_llm_judge_response( + raw_text, request.rubric, model, self.base_url, + ) + result.accumulators = {"elapsed_ms": elapsed_ms} + return result + + async def agrade(self, request: GradeRequest) -> GradeResponse: + """Grade via external OpenAI-compatible API (async).""" + model = request.model or self.model + body = self._build_request_body(request) + url = f"{self.base_url}/chat/completions" + + start_ms = time.time() * 1000 + + try: + async with httpx.AsyncClient(timeout=self.timeout) as client: + response = await client.post(url, json=body, headers=self._get_headers()) + response.raise_for_status() + except httpx.HTTPStatusError as e: + logger.warning("LLM API HTTP error: %s %s", e.response.status_code, e.response.text[:500]) + return GradeResponse( + normalized_score=0.0, + feedback=f"LLM API error: {e.response.status_code} - {e.response.text[:500]}", + model_used=model, + provider_used=self.base_url, + ) + except httpx.RequestError as e: + logger.warning("LLM API request failed: %s", e) + return GradeResponse( + normalized_score=0.0, + feedback=f"LLM API request failed: {e}", + model_used=model, + provider_used=self.base_url, + ) + + elapsed_ms = time.time() * 1000 - start_ms + data = response.json() + + # Extract response text from OpenAI format + raw_text = data["choices"][0]["message"]["content"] + + result = _parse_llm_judge_response( + raw_text, request.rubric, model, self.base_url, + ) + result.accumulators = {"elapsed_ms": elapsed_ms} + return result + + +# --------------------------------------------------------------------------- +# Auto-configuration from environment variables +# --------------------------------------------------------------------------- + + +def resolve_provider() -> Optional[LLMProvider]: + """Build an LLM provider from environment variables. + + Reads ``FLEET_LLM_*`` env vars and returns an ``ExternalProvider`` when + ``FLEET_LLM_API_KEY`` is set, otherwise returns ``None`` (meaning the + caller should fall back to the Fleet orchestrator). + + Env vars:: + + FLEET_LLM_API_KEY (required to activate external routing) + FLEET_LLM_BASE_URL (default: https://openrouter.ai/api/v1) + FLEET_LLM_MODEL (default: anthropic/claude-sonnet-4) + FLEET_LLM_TEMPERATURE (default: 0.0) + FLEET_LLM_MAX_TOKENS (default: 4096) + FLEET_LLM_TIMEOUT (default: 300) + + Returns: + An ``ExternalProvider`` if ``FLEET_LLM_API_KEY`` is set, else ``None``. + """ + api_key = os.environ.get(ENV_LLM_API_KEY) + if not api_key: + return None + + base_url = os.environ.get(ENV_LLM_BASE_URL) or None + model = os.environ.get(ENV_LLM_MODEL) or None + + temperature = 0.0 + temp_str = os.environ.get(ENV_LLM_TEMPERATURE) + if temp_str: + try: + temperature = float(temp_str) + except ValueError: + logger.warning("Invalid %s=%r, using default 0.0", ENV_LLM_TEMPERATURE, temp_str) + + max_tokens = 4096 + mt_str = os.environ.get(ENV_LLM_MAX_TOKENS) + if mt_str: + try: + max_tokens = int(mt_str) + except ValueError: + logger.warning("Invalid %s=%r, using default 4096", ENV_LLM_MAX_TOKENS, mt_str) + + timeout = 300.0 + to_str = os.environ.get(ENV_LLM_TIMEOUT) + if to_str: + try: + timeout = float(to_str) + except ValueError: + logger.warning("Invalid %s=%r, using default 300", ENV_LLM_TIMEOUT, to_str) + + logger.info( + "LLM provider configured from env: base_url=%s model=%s", + base_url or ExternalProvider.DEFAULT_BASE_URL, + model or ExternalProvider.DEFAULT_MODEL, + ) + + return ExternalProvider( + api_key=api_key, + base_url=base_url, + model=model, + temperature=temperature, + max_tokens=max_tokens, + timeout=timeout, + ) diff --git a/tests/test_llm_provider.py b/tests/test_llm_provider.py new file mode 100644 index 00000000..e5a7254c --- /dev/null +++ b/tests/test_llm_provider.py @@ -0,0 +1,1339 @@ +"""Tests for the LLM provider abstraction layer (fleet.llm_provider). + +Validates: +- FleetProvider delegates to the orchestrator client correctly +- ExternalProvider builds correct OpenAI-compatible requests +- ExternalProvider parses LLM JSON responses into GradeResponse +- SyncJudge routes through LLMProvider when configured +- Backward compatibility: SyncJudge still works without an LLMProvider +- resolve_provider() reads env vars and builds ExternalProvider +- Image.from_local / File.from_local read from local filesystem +""" + +import base64 +import json +import os +import tempfile +from typing import Any, Dict, Optional +from unittest.mock import MagicMock, patch, AsyncMock + +import httpx +import pytest + +from fleet.judge import ( + Criterion, + File, + Image, + JudgeResult, + Rubric, + SyncJudge, + _parse_grade_response, +) +from fleet.llm_provider import ( + ENV_LLM_API_KEY, + ENV_LLM_BASE_URL, + ENV_LLM_MODEL, + ENV_LLM_MAX_TOKENS, + ENV_LLM_TEMPERATURE, + ENV_LLM_TIMEOUT, + ExternalProvider, + FleetProvider, + GradeRequest, + GradeResponse, + LLMProvider, + resolve_provider, + _build_judge_user_message, + _parse_llm_judge_response, +) + + +# --------------------------------------------------------------------------- +# Helpers +# --------------------------------------------------------------------------- + +def _make_rubric() -> Rubric: + return Rubric(criteria=[ + Criterion(name="Accuracy", max=10, description="How accurate is the answer"), + Criterion(name="Clarity", max=5, description="How clear is the explanation"), + ]) + + +def _make_grade_request(rubric=None, submission="The answer is 42.") -> GradeRequest: + return GradeRequest( + rubric=rubric or _make_rubric(), + submission=submission, + ground_truth="42", + problem="What is the answer to life?", + context="This is a philosophy test.", + instance_id="test-instance-123", + ) + + +def _mock_orchestrator_response() -> dict: + """Response shape matching the Fleet orchestrator /v1/judge/grade.""" + return { + "normalized_score": 0.87, + "total_score": 13, + "max_score": 15, + "model_used": "claude-sonnet-4", + "provider_used": "anthropic", + "criteria": [ + {"name": "Accuracy", "score": 8, "max_score": 10, "reasoning": "Correct answer"}, + {"name": "Clarity", "score": 5, "max_score": 5, "reasoning": "Very clear"}, + ], + "feedback": "Good submission overall.", + } + + +def _mock_llm_json_response() -> str: + """JSON string mimicking what an external LLM returns.""" + return json.dumps({ + "criteria": [ + {"name": "Accuracy", "score": 8, "max_score": 10, "reasoning": "Correct answer"}, + {"name": "Clarity", "score": 5, "max_score": 5, "reasoning": "Very clear"}, + ], + "feedback": "Good submission overall.", + }) + + +def _clean_llm_env(monkeypatch): + """Remove all FLEET_LLM_* env vars for a clean test.""" + for var in [ENV_LLM_API_KEY, ENV_LLM_BASE_URL, ENV_LLM_MODEL, + ENV_LLM_TEMPERATURE, ENV_LLM_MAX_TOKENS, ENV_LLM_TIMEOUT]: + monkeypatch.delenv(var, raising=False) + + +# --------------------------------------------------------------------------- +# GradeResponse +# --------------------------------------------------------------------------- + + +class TestGradeResponse: + def test_to_dict_basic(self): + resp = GradeResponse( + normalized_score=0.8, + total_score=12, + max_score=15, + criteria=[{"name": "A", "score": 12, "max_score": 15, "reasoning": "ok"}], + feedback="Nice", + model_used="claude-sonnet-4", + provider_used="openrouter", + ) + d = resp.to_dict() + assert d["normalized_score"] == 0.8 + assert d["total_score"] == 12 + assert d["max_score"] == 15 + assert d["model_used"] == "claude-sonnet-4" + assert d["provider_used"] == "openrouter" + assert len(d["criteria"]) == 1 + + def test_to_dict_empty_criteria_excluded(self): + resp = GradeResponse(normalized_score=0.5) + d = resp.to_dict() + assert "criteria" not in d + assert "feedback" not in d + + +# --------------------------------------------------------------------------- +# _build_judge_user_message +# --------------------------------------------------------------------------- + + +class TestBuildJudgeUserMessage: + def test_includes_all_sections(self): + req = _make_grade_request() + msg = _build_judge_user_message(req) + assert "## Problem" in msg + assert "What is the answer to life?" in msg + assert "## Rubric" in msg + assert "Accuracy" in msg + assert "Clarity" in msg + assert "## Ground Truth" in msg + assert "42" in msg + assert "## Additional Context" in msg + assert "philosophy test" in msg + assert "## Submission to Grade" in msg + assert "The answer is 42." in msg + + def test_string_rubric(self): + req = GradeRequest(rubric="Grade from 1-10", submission="Hello") + msg = _build_judge_user_message(req) + assert "Grade from 1-10" in msg + + def test_no_submission(self): + req = GradeRequest(rubric="Test", submission=None) + msg = _build_judge_user_message(req) + assert "No submission text provided" in msg + + def test_conversation_included(self): + req = GradeRequest( + rubric="Test", + submission="Final answer", + conversation=[ + {"role": "user", "content": "What's 2+2?"}, + {"role": "assistant", "content": "4"}, + ], + ) + msg = _build_judge_user_message(req) + assert "Conversation History" in msg + assert "[user]: What's 2+2?" in msg + assert "[assistant]: 4" in msg + + +# --------------------------------------------------------------------------- +# _parse_llm_judge_response +# --------------------------------------------------------------------------- + + +class TestParseLLMJudgeResponse: + def test_valid_json(self): + rubric = _make_rubric() + resp = _parse_llm_judge_response( + _mock_llm_json_response(), rubric, "claude-sonnet-4", "openrouter" + ) + assert resp.normalized_score == pytest.approx(13 / 15, abs=0.01) + assert resp.total_score == 13 + assert resp.max_score == 15 + assert len(resp.criteria) == 2 + assert resp.feedback == "Good submission overall." + assert resp.model_used == "claude-sonnet-4" + assert resp.provider_used == "openrouter" + + def test_json_with_markdown_fences(self): + raw = f"```json\n{_mock_llm_json_response()}\n```" + resp = _parse_llm_judge_response(raw, _make_rubric(), "test", "test") + assert resp.normalized_score > 0 + + def test_invalid_json_returns_zero(self): + resp = _parse_llm_judge_response( + "This is not JSON", _make_rubric(), "test", "test" + ) + assert resp.normalized_score == 0.0 + assert "Failed to parse" in resp.feedback + + def test_string_rubric_max_from_criteria(self): + """When rubric is a plain string, max_score comes from criteria.""" + raw = json.dumps({ + "criteria": [ + {"name": "Quality", "score": 7, "max_score": 10, "reasoning": "Good"}, + ], + "feedback": "ok", + }) + resp = _parse_llm_judge_response(raw, "Grade it", "test", "test") + assert resp.max_score == 10 + assert resp.total_score == 7 + assert resp.normalized_score == pytest.approx(0.7, abs=0.01) + + +# --------------------------------------------------------------------------- +# FleetProvider +# --------------------------------------------------------------------------- + + +class TestFleetProvider: + def test_grade_delegates_to_client(self): + mock_client = MagicMock() + mock_response = MagicMock() + mock_response.json.return_value = _mock_orchestrator_response() + mock_client.request.return_value = mock_response + + provider = FleetProvider(client=mock_client, instance_id="inst-123") + req = _make_grade_request() + resp = provider.grade(req) + + # Verify the client was called + mock_client.request.assert_called_once() + call_args = mock_client.request.call_args + assert call_args[0] == ("POST", "/v1/judge/grade") + + # Verify the response + assert resp.normalized_score == pytest.approx(0.87, abs=0.01) + assert len(resp.criteria) == 2 + assert resp.model_used == "claude-sonnet-4" + + +# --------------------------------------------------------------------------- +# ExternalProvider +# --------------------------------------------------------------------------- + + +class TestExternalProvider: + def test_build_request_body(self): + provider = ExternalProvider( + api_key="sk-test", + model="anthropic/claude-sonnet-4", + ) + req = _make_grade_request() + body = provider._build_request_body(req) + + assert body["model"] == "anthropic/claude-sonnet-4" + assert body["temperature"] == 0.0 + assert len(body["messages"]) == 2 # system + user + assert body["messages"][0]["role"] == "system" + assert body["messages"][1]["role"] == "user" + + def test_model_override_from_request(self): + provider = ExternalProvider( + api_key="sk-test", + model="default/model", + ) + req = _make_grade_request() + req.model = "override/model" + body = provider._build_request_body(req) + assert body["model"] == "override/model" + + def test_grade_with_mocked_httpx(self): + provider = ExternalProvider( + api_key="sk-test", + base_url="https://test-api.example.com/v1", + model="test/model", + ) + + mock_response = MagicMock() + mock_response.json.return_value = { + "choices": [{ + "message": { + "content": _mock_llm_json_response(), + }, + }], + } + mock_response.raise_for_status = MagicMock() + + with patch("httpx.Client") as MockClient: + mock_client_instance = MagicMock() + mock_client_instance.post.return_value = mock_response + mock_client_instance.__enter__ = MagicMock(return_value=mock_client_instance) + mock_client_instance.__exit__ = MagicMock(return_value=False) + MockClient.return_value = mock_client_instance + + req = _make_grade_request() + resp = provider.grade(req) + + # Verify API call + mock_client_instance.post.assert_called_once() + call_args = mock_client_instance.post.call_args + assert call_args[0][0] == "https://test-api.example.com/v1/chat/completions" + + # Verify response parsing + assert resp.normalized_score == pytest.approx(13 / 15, abs=0.01) + assert len(resp.criteria) == 2 + assert resp.accumulators is not None + assert "elapsed_ms" in resp.accumulators + + def test_default_base_url(self): + provider = ExternalProvider(api_key="test") + assert provider.base_url == "https://openrouter.ai/api/v1" + + def test_custom_headers(self): + provider = ExternalProvider( + api_key="test", + extra_headers={"X-Custom": "value"}, + ) + headers = provider._get_headers() + assert headers["X-Custom"] == "value" + assert "Authorization" in headers + + +# --------------------------------------------------------------------------- +# Custom LLMProvider +# --------------------------------------------------------------------------- + + +class TestCustomProvider: + def test_custom_provider_implementation(self): + """Users can implement their own LLMProvider.""" + + class MyProvider(LLMProvider): + def grade(self, request: GradeRequest) -> GradeResponse: + return GradeResponse( + normalized_score=1.0, + total_score=10, + max_score=10, + criteria=[{"name": "Test", "score": 10, "max_score": 10, "reasoning": "Perfect"}], + feedback="Custom provider says: perfect!", + model_used="custom-model", + provider_used="my-provider", + ) + + provider = MyProvider() + req = _make_grade_request() + resp = provider.grade(req) + assert resp.normalized_score == 1.0 + assert resp.provider_used == "my-provider" + + +# --------------------------------------------------------------------------- +# resolve_provider() — env var auto-configuration +# --------------------------------------------------------------------------- + + +class TestResolveProvider: + def test_returns_none_when_no_api_key(self, monkeypatch): + """Without FLEET_LLM_API_KEY, resolve_provider returns None.""" + _clean_llm_env(monkeypatch) + assert resolve_provider() is None + + def test_returns_external_provider_with_api_key(self, monkeypatch): + """With FLEET_LLM_API_KEY set, returns ExternalProvider.""" + _clean_llm_env(monkeypatch) + monkeypatch.setenv(ENV_LLM_API_KEY, "sk-test-key-123") + provider = resolve_provider() + assert isinstance(provider, ExternalProvider) + assert provider.api_key == "sk-test-key-123" + # Defaults + assert provider.base_url == "https://openrouter.ai/api/v1" + assert provider.model == "anthropic/claude-sonnet-4" + assert provider.temperature == 0.0 + assert provider.max_tokens == 4096 + + def test_all_env_vars_respected(self, monkeypatch): + """All FLEET_LLM_* env vars are picked up.""" + _clean_llm_env(monkeypatch) + monkeypatch.setenv(ENV_LLM_API_KEY, "sk-custom") + monkeypatch.setenv(ENV_LLM_BASE_URL, "https://my-llm.example.com/v1") + monkeypatch.setenv(ENV_LLM_MODEL, "my-org/my-model") + monkeypatch.setenv(ENV_LLM_TEMPERATURE, "0.7") + monkeypatch.setenv(ENV_LLM_MAX_TOKENS, "8192") + monkeypatch.setenv(ENV_LLM_TIMEOUT, "600") + + provider = resolve_provider() + assert isinstance(provider, ExternalProvider) + assert provider.api_key == "sk-custom" + assert provider.base_url == "https://my-llm.example.com/v1" + assert provider.model == "my-org/my-model" + assert provider.temperature == pytest.approx(0.7) + assert provider.max_tokens == 8192 + assert provider.timeout == pytest.approx(600.0) + + def test_invalid_temperature_uses_default(self, monkeypatch): + _clean_llm_env(monkeypatch) + monkeypatch.setenv(ENV_LLM_API_KEY, "sk-test") + monkeypatch.setenv(ENV_LLM_TEMPERATURE, "not-a-number") + provider = resolve_provider() + assert provider.temperature == 0.0 + + def test_invalid_max_tokens_uses_default(self, monkeypatch): + _clean_llm_env(monkeypatch) + monkeypatch.setenv(ENV_LLM_API_KEY, "sk-test") + monkeypatch.setenv(ENV_LLM_MAX_TOKENS, "banana") + provider = resolve_provider() + assert provider.max_tokens == 4096 + + +# --------------------------------------------------------------------------- +# SyncJudge auto-resolve from env vars +# --------------------------------------------------------------------------- + + +class TestSyncJudgeEnvAutoResolve: + def test_auto_resolves_from_env(self, monkeypatch): + """SyncJudge auto-detects FLEET_LLM_API_KEY and routes externally.""" + _clean_llm_env(monkeypatch) + monkeypatch.setenv(ENV_LLM_API_KEY, "sk-auto-test") + monkeypatch.setenv(ENV_LLM_MODEL, "test/auto-model") + + judge = SyncJudge(client=None, instance_id="auto-test") + assert judge._llm_provider is not None + assert isinstance(judge._llm_provider, ExternalProvider) + assert judge._llm_provider.model == "test/auto-model" + + def test_defaults_to_fleet_when_no_env(self, monkeypatch): + """Without env vars, SyncJudge._llm_provider is None (Fleet route).""" + _clean_llm_env(monkeypatch) + mock_client = MagicMock() + judge = SyncJudge(client=mock_client, instance_id="fleet-test") + assert judge._llm_provider is None + + def test_explicit_provider_overrides_env(self, monkeypatch): + """Explicit llm_provider kwarg takes priority over env vars.""" + _clean_llm_env(monkeypatch) + monkeypatch.setenv(ENV_LLM_API_KEY, "sk-from-env") + + class CustomProv(LLMProvider): + def grade(self, request): + return GradeResponse(normalized_score=0.42) + + custom = CustomProv() + judge = SyncJudge(client=None, instance_id="test", llm_provider=custom) + assert judge._llm_provider is custom + + def test_explicit_none_disables_env_auto(self, monkeypatch): + """Passing llm_provider=None explicitly skips env auto-detect.""" + _clean_llm_env(monkeypatch) + monkeypatch.setenv(ENV_LLM_API_KEY, "sk-should-be-ignored") + + mock_client = MagicMock() + judge = SyncJudge(client=mock_client, instance_id="test", llm_provider=None) + assert judge._llm_provider is None + + +# --------------------------------------------------------------------------- +# SyncJudge integration (from earlier tests) +# --------------------------------------------------------------------------- + + +class TestSyncJudgeWithProvider: + def test_routes_through_llm_provider(self, monkeypatch): + """When llm_provider is set, grade() uses it instead of the client.""" + _clean_llm_env(monkeypatch) + + class StubProvider(LLMProvider): + def __init__(self): + self.called = False + + def grade(self, request: GradeRequest) -> GradeResponse: + self.called = True + return GradeResponse( + normalized_score=0.95, + total_score=19, + max_score=20, + criteria=[ + {"name": "Accuracy", "score": 10, "max_score": 10, "reasoning": "Spot on"}, + {"name": "Clarity", "score": 9, "max_score": 10, "reasoning": "Almost perfect"}, + ], + feedback="Excellent work", + model_used="stub-model", + provider_used="stub", + ) + + stub = StubProvider() + judge = SyncJudge(client=None, instance_id="local-123", llm_provider=stub) + result = judge.grade( + _make_rubric(), + "The answer is 42.", + ground_truth="42", + problem="What is the meaning of life?", + ) + + assert stub.called + assert isinstance(result, JudgeResult) + assert float(result) == pytest.approx(0.95, abs=0.01) + assert result.criteria is not None + assert len(result.criteria) == 2 + + def test_default_routes_through_client(self, monkeypatch): + """Without llm_provider, grade() uses the orchestrator client.""" + _clean_llm_env(monkeypatch) + mock_client = MagicMock() + mock_response = MagicMock() + mock_response.json.return_value = _mock_orchestrator_response() + mock_client.request.return_value = mock_response + + judge = SyncJudge(client=mock_client, instance_id="inst-456") + result = judge.grade(_make_rubric(), "Answer") + + mock_client.request.assert_called_once() + assert float(result) == pytest.approx(0.87, abs=0.01) + + def test_reference_claims_folded_into_context(self, monkeypatch): + """reference_claims should be folded into context for both paths.""" + _clean_llm_env(monkeypatch) + + class CapturingProvider(LLMProvider): + def __init__(self): + self.last_request = None + + def grade(self, request: GradeRequest) -> GradeResponse: + self.last_request = request + return GradeResponse(normalized_score=0.5) + + provider = CapturingProvider() + judge = SyncJudge(client=None, instance_id="local", llm_provider=provider) + + judge.grade( + "Simple rubric", + "submission", + context="Some context", + reference_claims="Claim 1, Claim 2", + ) + + assert provider.last_request is not None + assert "Some context" in provider.last_request.context + assert "Reference Claims" in provider.last_request.context + assert "Claim 1, Claim 2" in provider.last_request.context + + +# --------------------------------------------------------------------------- +# Image.from_local / File.from_local +# --------------------------------------------------------------------------- + + +class TestImageFromLocal: + def test_from_local_creates_local_source(self): + img = Image.from_local("/tmp/test.png") + assert img.source == "local" + assert img._local_path == "/tmp/test.png" + assert img.filename == "test.png" + assert img.media_type == "image/png" + + def test_from_local_serializes_to_base64(self): + """from_local reads the file and serializes as base64.""" + raw_bytes = b"\x89PNG\r\n\x1a\nfake-png-content" + with tempfile.NamedTemporaryFile(suffix=".png", delete=False) as f: + f.write(raw_bytes) + path = f.name + try: + img = Image.from_local(path) + d = img.serialize() + assert d["source"] == "base64" + assert d["media_type"] == "image/png" + # Verify the data decodes back to original bytes + assert base64.b64decode(d["data"]) == raw_bytes + finally: + os.unlink(path) + + def test_from_local_raises_on_missing_file(self): + img = Image.from_local("/tmp/nonexistent_image_12345.png") + with pytest.raises(ValueError, match="Cannot read local image"): + img.serialize() + + def test_from_local_media_type_override(self): + img = Image.from_local("/tmp/photo.webp", media_type="image/webp") + assert img.media_type == "image/webp" + + def test_from_local_guesses_jpeg(self): + img = Image.from_local("/tmp/photo.jpg") + assert img.media_type == "image/jpeg" + + def test_s3_still_works(self): + """Ensure S3 constructor is not broken.""" + img = Image.s3("s3://bucket/key.png", media_type="image/png") + assert img.source == "s3" + d = img.serialize() + assert d["source"] == "s3" + assert d["url"] == "s3://bucket/key.png" + + def test_from_url_still_works(self): + """Ensure URL constructor is not broken.""" + img = Image.from_url("https://example.com/img.png") + assert img.source == "url" + d = img.serialize() + assert d["source"] == "url" + assert d["url"] == "https://example.com/img.png" + + +class TestFileFromLocal: + def test_from_local_creates_local_source(self): + f = File.from_local("/tmp/report.pdf") + assert f.source == "local" + assert f._local_path == "/tmp/report.pdf" + assert f.filename == "report.pdf" + assert f.media_type == "application/pdf" + + def test_from_local_serializes_to_base64(self): + raw_bytes = b"%PDF-1.4 fake pdf content" + with tempfile.NamedTemporaryFile(suffix=".pdf", delete=False) as tf: + tf.write(raw_bytes) + path = tf.name + try: + f = File.from_local(path) + d = f.serialize() + assert d["source"] == "base64" + assert d["media_type"] == "application/pdf" + assert d["filename"] == os.path.basename(path) + assert base64.b64decode(d["data"]) == raw_bytes + finally: + os.unlink(path) + + def test_from_local_raises_on_missing_file(self): + f = File.from_local("/tmp/nonexistent_file_12345.pdf") + with pytest.raises(ValueError, match="Cannot read local file"): + f.serialize() + + def test_from_local_csv(self): + raw_bytes = b"name,value\nalice,42\n" + with tempfile.NamedTemporaryFile(suffix=".csv", delete=False) as tf: + tf.write(raw_bytes) + path = tf.name + try: + f = File.from_local(path) + assert f.media_type == "text/csv" + d = f.serialize() + assert base64.b64decode(d["data"]) == raw_bytes + finally: + os.unlink(path) + + def test_s3_still_works(self): + """Ensure S3 constructor is not broken.""" + f = File.s3("s3://bucket/data.csv", media_type="text/csv") + assert f.source == "s3" + d = f.serialize() + assert d["source"] == "s3" + assert d["url"] == "s3://bucket/data.csv" + + +# --------------------------------------------------------------------------- +# ExternalProvider with local images +# --------------------------------------------------------------------------- + + +class TestExternalProviderLocalImages: + def test_local_image_resolved_in_request_body(self): + """Local images are resolved to base64 in the LLM request body.""" + raw_bytes = b"\x89PNG\r\n\x1a\nfake-png" + with tempfile.NamedTemporaryFile(suffix=".png", delete=False) as f: + f.write(raw_bytes) + path = f.name + try: + provider = ExternalProvider(api_key="sk-test", model="test/model") + img = Image.from_local(path) + req = GradeRequest( + rubric="Test rubric", + submission="Answer", + images={"screenshot": img}, + ) + body = provider._build_request_body(req) + + # The user message should contain an image_url block with data: URI + user_msg = body["messages"][1] + content_blocks = user_msg["content"] + image_blocks = [b for b in content_blocks if b.get("type") == "image_url"] + assert len(image_blocks) == 1 + assert image_blocks[0]["image_url"]["url"].startswith("data:image/png;base64,") + finally: + os.unlink(path) + + +# --------------------------------------------------------------------------- +# End-to-end: ExternalProvider + SyncJudge +# --------------------------------------------------------------------------- + + +class TestEndToEnd: + def test_external_provider_with_sync_judge(self, monkeypatch): + """Full integration: ExternalProvider → SyncJudge → JudgeResult.""" + _clean_llm_env(monkeypatch) + provider = ExternalProvider( + api_key="sk-test", + model="test/model", + ) + + mock_response = MagicMock() + mock_response.json.return_value = { + "choices": [{ + "message": { + "content": _mock_llm_json_response(), + }, + }], + } + mock_response.raise_for_status = MagicMock() + + with patch("httpx.Client") as MockClient: + mock_client_instance = MagicMock() + mock_client_instance.post.return_value = mock_response + mock_client_instance.__enter__ = MagicMock(return_value=mock_client_instance) + mock_client_instance.__exit__ = MagicMock(return_value=False) + MockClient.return_value = mock_client_instance + + judge = SyncJudge(client=None, instance_id="local-e2e", llm_provider=provider) + result = judge.grade( + _make_rubric(), + "The answer is 42.", + ground_truth="42", + ) + + assert isinstance(result, JudgeResult) + assert isinstance(result, float) + assert float(result) > 0 + assert result.criteria is not None + assert len(result.criteria) == 2 + + def test_env_var_auto_config_e2e(self, monkeypatch): + """Full e2e: env vars → auto ExternalProvider → SyncJudge → JudgeResult.""" + _clean_llm_env(monkeypatch) + monkeypatch.setenv(ENV_LLM_API_KEY, "sk-auto-e2e") + monkeypatch.setenv(ENV_LLM_BASE_URL, "https://e2e-api.example.com/v1") + monkeypatch.setenv(ENV_LLM_MODEL, "test/e2e-model") + + mock_response = MagicMock() + mock_response.json.return_value = { + "choices": [{ + "message": { + "content": _mock_llm_json_response(), + }, + }], + } + mock_response.raise_for_status = MagicMock() + + with patch("httpx.Client") as MockClient: + mock_client_instance = MagicMock() + mock_client_instance.post.return_value = mock_response + mock_client_instance.__enter__ = MagicMock(return_value=mock_client_instance) + mock_client_instance.__exit__ = MagicMock(return_value=False) + MockClient.return_value = mock_client_instance + + # No explicit provider — should auto-detect from env + judge = SyncJudge(client=None, instance_id="env-e2e") + assert isinstance(judge._llm_provider, ExternalProvider) + + result = judge.grade(_make_rubric(), "The answer is 42.") + + # Verify it called the right URL + call_args = mock_client_instance.post.call_args + assert call_args[0][0] == "https://e2e-api.example.com/v1/chat/completions" + + assert isinstance(result, JudgeResult) + assert float(result) > 0 + + +# --------------------------------------------------------------------------- +# Image.from_path — source-agnostic constructor +# --------------------------------------------------------------------------- + + +class TestImageFromPath: + def test_creates_path_source(self): + img = Image.from_path("screenshots/gold.png") + assert img.source == "path" + assert img._path == "screenshots/gold.png" + assert img.filename == "gold.png" + assert img.media_type == "image/png" + + def test_s3_uri(self): + img = Image.from_path("s3://bucket/screenshots/gold.png") + assert img.source == "path" + assert img._path == "s3://bucket/screenshots/gold.png" + + def test_http_url(self): + img = Image.from_path("https://example.com/gold.png") + assert img.source == "path" + assert img._path == "https://example.com/gold.png" + + def test_absolute_local(self): + img = Image.from_path("/data/images/gold.png") + assert img.source == "path" + assert img._path == "/data/images/gold.png" + + def test_media_type_override(self): + img = Image.from_path("image.webp", media_type="image/webp") + assert img.media_type == "image/webp" + + def test_serialize_s3_fallback(self): + """When no provider resolves, serialize auto-detects s3:// scheme.""" + img = Image.from_path("s3://bucket/key.png") + d = img.serialize() + assert d["source"] == "s3" + assert d["url"] == "s3://bucket/key.png" + + def test_serialize_http_fallback(self): + """When no provider resolves, serialize auto-detects https:// scheme.""" + img = Image.from_path("https://example.com/img.png") + d = img.serialize() + assert d["source"] == "url" + assert d["url"] == "https://example.com/img.png" + + def test_serialize_local_fallback(self): + """When no provider resolves, serialize reads local file.""" + raw_bytes = b"\x89PNG\r\n\x1a\nfake-png-content" + with tempfile.NamedTemporaryFile(suffix=".png", delete=False) as f: + f.write(raw_bytes) + path = f.name + try: + img = Image.from_path(path) + d = img.serialize() + assert d["source"] == "base64" + assert d["media_type"] == "image/png" + assert base64.b64decode(d["data"]) == raw_bytes + finally: + os.unlink(path) + + def test_serialize_missing_local_raises(self): + img = Image.from_path("/nonexistent/path/img.png") + with pytest.raises(ValueError, match="Cannot read image path"): + img.serialize() + + +# --------------------------------------------------------------------------- +# File.from_path — source-agnostic constructor +# --------------------------------------------------------------------------- + + +class TestFileFromPath: + def test_creates_path_source(self): + f = File.from_path("reports/output.pdf") + assert f.source == "path" + assert f._path == "reports/output.pdf" + assert f.filename == "output.pdf" + assert f.media_type == "application/pdf" + + def test_s3_uri(self): + f = File.from_path("s3://bucket/data.csv") + assert f.source == "path" + assert f._path == "s3://bucket/data.csv" + + def test_serialize_s3_fallback(self): + f = File.from_path("s3://bucket/data.csv") + d = f.serialize() + assert d["source"] == "s3" + assert d["url"] == "s3://bucket/data.csv" + + def test_serialize_local_fallback(self): + raw_bytes = b"name,value\nalice,42\n" + with tempfile.NamedTemporaryFile(suffix=".csv", delete=False) as tf: + tf.write(raw_bytes) + path = tf.name + try: + f = File.from_path(path) + d = f.serialize() + assert d["source"] == "base64" + assert d["media_type"] == "text/csv" + assert base64.b64decode(d["data"]) == raw_bytes + finally: + os.unlink(path) + + def test_serialize_missing_local_raises(self): + f = File.from_path("/nonexistent/path/file.pdf") + with pytest.raises(ValueError, match="Cannot read file path"): + f.serialize() + + +# --------------------------------------------------------------------------- +# LLMProvider.resolve_image / resolve_file +# --------------------------------------------------------------------------- + + +class TestProviderResolve: + def test_resolve_image_s3(self): + """resolve_image converts s3:// path to Image.s3().""" + provider = ExternalProvider(api_key="sk-test") + img = Image.from_path("s3://bucket/key.png") + resolved = provider.resolve_image(img) + assert resolved.source == "s3" + assert resolved.url == "s3://bucket/key.png" + + def test_resolve_image_http(self): + """resolve_image converts https:// path to Image.from_url().""" + provider = ExternalProvider(api_key="sk-test") + img = Image.from_path("https://example.com/img.png") + resolved = provider.resolve_image(img) + assert resolved.source == "url" + assert resolved.url == "https://example.com/img.png" + + def test_resolve_image_local(self): + """resolve_image converts bare path to Image.from_local().""" + provider = ExternalProvider(api_key="sk-test") + img = Image.from_path("/data/images/gold.png") + resolved = provider.resolve_image(img) + assert resolved.source == "local" + assert resolved._local_path == "/data/images/gold.png" + + def test_resolve_image_noop_for_non_path(self): + """resolve_image passes through non-path images unchanged.""" + provider = ExternalProvider(api_key="sk-test") + img = Image.s3("s3://bucket/key.png") + resolved = provider.resolve_image(img) + assert resolved is img # same object + + def test_resolve_file_s3(self): + provider = ExternalProvider(api_key="sk-test") + f = File.from_path("s3://bucket/data.csv") + resolved = provider.resolve_file(f) + assert resolved.source == "s3" + assert resolved.url == "s3://bucket/data.csv" + + def test_resolve_file_local(self): + provider = ExternalProvider(api_key="sk-test") + f = File.from_path("/data/reports/output.pdf") + resolved = provider.resolve_file(f) + assert resolved.source == "local" + assert resolved._local_path == "/data/reports/output.pdf" + + def test_resolve_images_dict(self): + provider = ExternalProvider(api_key="sk-test") + images = { + "gold": Image.from_path("s3://bucket/gold.png"), + "agent": Image.from_path("/local/agent.png"), + "ref": Image.s3("s3://bucket/ref.png"), # non-path, should pass through + } + resolved = provider.resolve_images(images) + assert resolved["gold"].source == "s3" + assert resolved["agent"].source == "local" + assert resolved["ref"].source == "s3" + assert resolved["ref"] is images["ref"] # unchanged + + def test_resolve_images_none(self): + provider = ExternalProvider(api_key="sk-test") + assert provider.resolve_images(None) is None + + def test_resolve_files_dict(self): + provider = ExternalProvider(api_key="sk-test") + files = { + "report": File.from_path("s3://bucket/report.pdf"), + "local": File.from_path("/data/local.csv"), + } + resolved = provider.resolve_files(files) + assert resolved["report"].source == "s3" + assert resolved["local"].source == "local" + + +# --------------------------------------------------------------------------- +# Custom provider with resolve override +# --------------------------------------------------------------------------- + + +class TestCustomProviderResolve: + def test_custom_resolve_prepends_s3_prefix(self, monkeypatch): + """Custom provider can override resolve to prepend S3 prefix.""" + _clean_llm_env(monkeypatch) + + class S3PrefixProvider(LLMProvider): + """Provider that prepends an S3 bucket prefix to bare paths.""" + + def __init__(self, bucket: str): + self.bucket = bucket + + def resolve_image(self, image): + if getattr(image, "source", None) != "path": + return image + path = image._path or image.filename or "" + if not path.startswith(("s3://", "http://", "https://")): + # Prepend S3 bucket prefix + s3_url = f"s3://{self.bucket}/{path}" + return Image.s3(s3_url, media_type=image.media_type) + return super().resolve_image(image) + + def grade(self, request): + return GradeResponse(normalized_score=1.0) + + provider = S3PrefixProvider(bucket="my-images-bucket") + + # Bare path → gets s3 prefix + img = Image.from_path("screenshots/gold.png") + resolved = provider.resolve_image(img) + assert resolved.source == "s3" + assert resolved.url == "s3://my-images-bucket/screenshots/gold.png" + + # Already has s3:// → passed through normally + img2 = Image.from_path("s3://other-bucket/img.png") + resolved2 = provider.resolve_image(img2) + assert resolved2.source == "s3" + assert resolved2.url == "s3://other-bucket/img.png" + + def test_judge_calls_resolve_before_grade(self, monkeypatch): + """SyncJudge calls provider.resolve_images before grade.""" + _clean_llm_env(monkeypatch) + + class TrackingProvider(LLMProvider): + def __init__(self): + self.resolved_images = None + + def resolve_image(self, image): + if getattr(image, "source", None) != "path": + return image + # Convert all paths to base64 with marker data + return Image.from_base64("RESOLVED", image.filename or "img.png") + + def grade(self, request): + self.resolved_images = request.images + return GradeResponse(normalized_score=1.0) + + provider = TrackingProvider() + judge = SyncJudge(client=None, instance_id="test", llm_provider=provider) + judge.grade( + "test rubric", + "submission", + images={"gold": Image.from_path("gold.png")}, + ) + + assert provider.resolved_images is not None + assert provider.resolved_images["gold"].source == "base64" + assert provider.resolved_images["gold"].data == "RESOLVED" + + +# --------------------------------------------------------------------------- +# ExternalProvider with from_path images in request body +# --------------------------------------------------------------------------- + + +class TestExternalProviderPathImages: + def test_path_image_resolved_in_request_body(self): + """from_path images are resolved to base64 in the LLM request body.""" + raw_bytes = b"\x89PNG\r\n\x1a\nfake-png" + with tempfile.NamedTemporaryFile(suffix=".png", delete=False) as f: + f.write(raw_bytes) + path = f.name + try: + provider = ExternalProvider(api_key="sk-test", model="test/model") + img = Image.from_path(path) + req = GradeRequest( + rubric="Test rubric", + submission="Answer", + images={"screenshot": img}, + ) + body = provider._build_request_body(req) + + user_msg = body["messages"][1] + content_blocks = user_msg["content"] + image_blocks = [b for b in content_blocks if b.get("type") == "image_url"] + assert len(image_blocks) == 1 + assert image_blocks[0]["image_url"]["url"].startswith("data:image/png;base64,") + finally: + os.unlink(path) + + +# --------------------------------------------------------------------------- +# ExternalProvider file handling in request body +# --------------------------------------------------------------------------- + + +class TestExternalProviderFiles: + def test_text_file_included_in_request_body(self): + """Text files are inlined as text blocks in the LLM request.""" + with tempfile.NamedTemporaryFile(suffix=".csv", mode="w", delete=False) as f: + f.write("name,score\nAlice,95\nBob,87\n") + path = f.name + try: + provider = ExternalProvider(api_key="sk-test", model="test/model") + req = GradeRequest( + rubric="Test rubric", + submission="Answer", + files={"data": File.from_local(path, media_type="text/csv")}, + ) + body = provider._build_request_body(req) + + user_msg = body["messages"][1] + content_blocks = user_msg["content"] + text_blocks = [b for b in content_blocks if b.get("type") == "text"] + # Should have a file block + the main user message + file_blocks = [b for b in text_blocks if "## File:" in b.get("text", "")] + assert len(file_blocks) == 1 + assert "Alice,95" in file_blocks[0]["text"] + finally: + os.unlink(path) + + def test_base64_file_included_in_request_body(self): + """Base64-encoded files are decoded and inlined as text.""" + content = "Hello, world!" + b64_data = base64.b64encode(content.encode()).decode() + provider = ExternalProvider(api_key="sk-test", model="test/model") + req = GradeRequest( + rubric="Test rubric", + submission="Answer", + files={"readme": File.from_base64(b64_data, "readme.txt", media_type="text/plain")}, + ) + body = provider._build_request_body(req) + + user_msg = body["messages"][1] + text_blocks = [b for b in user_msg["content"] if b.get("type") == "text"] + file_blocks = [b for b in text_blocks if "## File:" in b.get("text", "")] + assert len(file_blocks) == 1 + assert "Hello, world!" in file_blocks[0]["text"] + + def test_path_file_included_in_request_body(self): + """from_path files are read from disk and inlined.""" + with tempfile.NamedTemporaryFile(suffix=".json", mode="w", delete=False) as f: + json.dump({"key": "value"}, f) + path = f.name + try: + provider = ExternalProvider(api_key="sk-test", model="test/model") + req = GradeRequest( + rubric="Test rubric", + submission="Answer", + files={"config": File.from_path(path)}, + ) + body = provider._build_request_body(req) + + user_msg = body["messages"][1] + text_blocks = [b for b in user_msg["content"] if b.get("type") == "text"] + file_blocks = [b for b in text_blocks if "## File:" in b.get("text", "")] + assert len(file_blocks) == 1 + assert '"key"' in file_blocks[0]["text"] + finally: + os.unlink(path) + + def test_unresolvable_file_skipped(self): + """Files that can't be resolved are skipped with a warning.""" + provider = ExternalProvider(api_key="sk-test", model="test/model") + req = GradeRequest( + rubric="Test rubric", + submission="Answer", + files={"missing": File.from_local("/nonexistent/file.csv")}, + ) + body = provider._build_request_body(req) + + user_msg = body["messages"][1] + text_blocks = [b for b in user_msg["content"] if b.get("type") == "text"] + file_blocks = [b for b in text_blocks if "## File:" in b.get("text", "")] + assert len(file_blocks) == 0 + + +# --------------------------------------------------------------------------- +# ExternalProvider error handling +# --------------------------------------------------------------------------- + + +class TestExternalProviderErrorHandling: + def test_http_error_returns_zero_grade(self): + """HTTP errors return GradeResponse with 0.0 score instead of raising.""" + provider = ExternalProvider(api_key="sk-test", model="test/model") + + mock_response = MagicMock() + mock_response.status_code = 500 + mock_response.text = "Internal Server Error" + mock_response.raise_for_status.side_effect = httpx.HTTPStatusError( + "Server Error", request=MagicMock(), response=mock_response, + ) + + with patch("httpx.Client") as MockClient: + mock_client = MagicMock() + mock_client.post.return_value = mock_response + mock_client.__enter__ = MagicMock(return_value=mock_client) + mock_client.__exit__ = MagicMock(return_value=False) + MockClient.return_value = mock_client + + req = _make_grade_request() + resp = provider.grade(req) + + assert resp.normalized_score == 0.0 + assert "LLM API error" in resp.feedback + assert "500" in resp.feedback + + def test_request_error_returns_zero_grade(self): + """Network errors return GradeResponse with 0.0 score instead of raising.""" + provider = ExternalProvider(api_key="sk-test", model="test/model") + + with patch("httpx.Client") as MockClient: + mock_client = MagicMock() + mock_client.post.side_effect = httpx.RequestError("Connection refused") + mock_client.__enter__ = MagicMock(return_value=mock_client) + mock_client.__exit__ = MagicMock(return_value=False) + MockClient.return_value = mock_client + + req = _make_grade_request() + resp = provider.grade(req) + + assert resp.normalized_score == 0.0 + assert "request failed" in resp.feedback.lower() + + +# --------------------------------------------------------------------------- +# ExternalProvider api_key validation +# --------------------------------------------------------------------------- + + +class TestExternalProviderValidation: + def test_empty_api_key_raises(self): + with pytest.raises(ValueError, match="non-empty"): + ExternalProvider(api_key="") + + def test_whitespace_api_key_raises(self): + with pytest.raises(ValueError, match="non-empty"): + ExternalProvider(api_key=" ") + + +# --------------------------------------------------------------------------- +# Async: ExternalProvider.agrade() +# --------------------------------------------------------------------------- + + +class TestExternalProviderAsync: + @pytest.mark.asyncio + async def test_agrade_with_mocked_httpx(self): + """ExternalProvider.agrade() uses httpx.AsyncClient correctly.""" + provider = ExternalProvider( + api_key="sk-test", + base_url="https://test-api.example.com/v1", + model="test/model", + ) + + mock_response = MagicMock() + mock_response.json.return_value = { + "choices": [{"message": {"content": _mock_llm_json_response()}}], + } + mock_response.raise_for_status = MagicMock() + + with patch("fleet.llm_provider.httpx.AsyncClient") as MockAsyncClient: + mock_client = AsyncMock() + mock_client.post.return_value = mock_response + mock_client.__aenter__ = AsyncMock(return_value=mock_client) + mock_client.__aexit__ = AsyncMock(return_value=False) + MockAsyncClient.return_value = mock_client + + req = _make_grade_request() + resp = await provider.agrade(req) + + mock_client.post.assert_called_once() + assert resp.normalized_score == pytest.approx(13 / 15, abs=0.01) + assert len(resp.criteria) == 2 + + @pytest.mark.asyncio + async def test_agrade_http_error_returns_zero(self): + """Async HTTP errors return GradeResponse with 0.0 score.""" + provider = ExternalProvider(api_key="sk-test", model="test/model") + + mock_response = MagicMock() + mock_response.status_code = 429 + mock_response.text = "Rate limited" + mock_response.raise_for_status.side_effect = httpx.HTTPStatusError( + "Rate limited", request=MagicMock(), response=mock_response, + ) + + with patch("fleet.llm_provider.httpx.AsyncClient") as MockAsyncClient: + mock_client = AsyncMock() + mock_client.post.return_value = mock_response + mock_client.__aenter__ = AsyncMock(return_value=mock_client) + mock_client.__aexit__ = AsyncMock(return_value=False) + MockAsyncClient.return_value = mock_client + + req = _make_grade_request() + resp = await provider.agrade(req) + + assert resp.normalized_score == 0.0 + assert "429" in resp.feedback + + +# --------------------------------------------------------------------------- +# Async: LLMProvider.agrade() default (run_in_executor) +# --------------------------------------------------------------------------- + + +class TestLLMProviderAgradeDefault: + @pytest.mark.asyncio + async def test_default_agrade_calls_grade_without_blocking(self): + """Base class agrade() offloads to executor so it doesn't block.""" + + class SyncOnlyProvider(LLMProvider): + def grade(self, request: GradeRequest) -> GradeResponse: + return GradeResponse( + normalized_score=0.75, + model_used="sync-model", + provider_used="sync-provider", + ) + + provider = SyncOnlyProvider() + req = _make_grade_request() + resp = await provider.agrade(req) + + assert resp.normalized_score == 0.75 + assert resp.model_used == "sync-model" + + +# --------------------------------------------------------------------------- +# Async: AsyncJudge + provider integration +# --------------------------------------------------------------------------- + + +class TestAsyncJudgeWithProvider: + @pytest.mark.asyncio + async def test_async_judge_routes_through_provider(self): + """AsyncJudge.grade() calls provider.agrade() when provider is set.""" + from fleet._async.judge import AsyncJudge + + mock_provider = AsyncMock(spec=LLMProvider) + mock_provider.agrade.return_value = GradeResponse( + normalized_score=0.9, + total_score=9, + max_score=10, + criteria=[{"name": "Test", "score": 9, "max_score": 10, "reasoning": "Good"}], + feedback="Nice work", + model_used="test-model", + provider_used="test-provider", + ) + mock_provider.resolve_images.return_value = None + mock_provider.resolve_files.return_value = None + + judge = AsyncJudge(client=None, instance_id="test", llm_provider=mock_provider) + result = await judge.grade("Test rubric", "My submission") + + mock_provider.agrade.assert_called_once() + assert float(result) == pytest.approx(0.9) diff --git a/uv.lock b/uv.lock index db8a2922..141c5903 100644 --- a/uv.lock +++ b/uv.lock @@ -1,5 +1,5 @@ version = 1 -revision = 2 +revision = 3 requires-python = ">=3.9" resolution-markers = [ "python_full_version >= '3.10'",