diff --git a/.github/workflows/ci.yml b/.github/workflows/ci.yml index 0fd3258..421d53b 100644 --- a/.github/workflows/ci.yml +++ b/.github/workflows/ci.yml @@ -35,7 +35,6 @@ jobs: run: python -m black --check src/prkit tests/prkit - name: Type check - continue-on-error: true run: python -m mypy src/prkit - name: Test package diff --git a/CHANGELOG.md b/CHANGELOG.md index a33a757..b53603a 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -9,6 +9,10 @@ Production releases follow semantic versioning. TestPyPI validation builds use P ### Added +- **`OpenAIModel` custom endpoint support** — new keyword-only constructor params `base_url`, `api_key`, and `api_key_env` allow routing to any proxy or gateway that implements the OpenAI Responses API (`POST /v1/responses`) with an explicit key or key from a named environment variable. Backward-compatible: omitting all three preserves existing `OPENAI_API_KEY` + default endpoint behaviour. +- **`OllamaModel` explicit auth params** — new keyword-only constructor params `api_key` and `api_key_env` forward a `Bearer` token as the `Authorization` header to `ollama.Client`, providing API-key parity with other providers. Works for cloud endpoints (e.g. `base_url="https://ollama.com"`). +- **Remote-safe Ollama preflight** — when `base_url` or `OLLAMA_HOST` points to a non-local host, a failed startup connectivity check now emits a warning instead of raising `ConnectionError`; precise errors surface at `chat()` call time. +- **"Extending prkit" contract documented** — `DATASETS.md` and `CORE.md` now document the stable external extension points: registering a custom `DatasetHub` loader/downloader from outside the package, local-directory loading without a downloader, `OpenAIModel` / `OllamaModel` custom-endpoint construction, and `register_model_client` for additional providers. - **`prkit` command-line interface** (`prkit list`, `prkit info `, `prkit download `, `prkit --version`) for dataset workflows, installed via the `prkit` console script. - PEP 561 typing support: ships a `py.typed` marker and the `Typing :: Typed` classifier. - `ruff` linting + import sorting, a `.pre-commit-config.yaml`, a `Makefile`, and GitHub Actions CI (lint, format check, type check, tests on Python 3.10–3.12) plus a release workflow. @@ -30,8 +34,13 @@ Production releases follow semantic versioning. TestPyPI validation builds use P - Coverage enforcement for `prkit` now uses a 60% minimum and keeps unit tests in pytest format. - Provider-model test targets were updated for OpenAI, Gemini, Anthropic, Ollama, DeepSeek, xAI, and DashScope clients. +### Changed + +- **(internal)** Model-output JSON extraction consolidated: the duplicate `extract_json_object` in `prkit.evaluation.llm_judge.parse` and the unreachable helpers `_iter_braced_json_candidates`, `_try_parse_json_object`, `_JSON_FENCE_RE`, and the thin `_extract_json_object` wrapper in `prkit.semantics.inference.calls` are removed. All call sites now delegate to the single canonical `extract_json_object` / `extract_json_payload` in `prkit.core.model_clients.structured_output`. Public API and parsing semantics are unchanged. + ### Fixed +- **`DatasetHub` registration-ordering bug** — calling `DatasetHub.register(name, Loader)` before any built-in dataset was touched caused all built-in loaders and downloaders to be permanently suppressed. Built-ins are now seeded idempotently (via `setdefault`) at the start of every public mutating method, so external registrations can safely happen in any order. - JEEBench loader handling for numeric answer categories and retained metadata. - Workflow module behavior in domain assessment, theorem review, and workflow composition paths. diff --git a/CORE.md b/CORE.md index 38230a8..147afa5 100644 --- a/CORE.md +++ b/CORE.md @@ -195,6 +195,81 @@ text = client.chat( print(text) ``` +#### Custom OpenAI Responses-API endpoints + +`OpenAIModel` accepts `base_url` and `api_key` / `api_key_env` keyword arguments for +routing to a proxy or gateway that implements the OpenAI **Responses API** +(`POST /v1/responses`). These are **not** available through `create_model_client` (which +is routing-only); construct `OpenAIModel` directly: + +```python +from prkit.core.model_clients.openai import OpenAIModel + +# Explicit key + custom endpoint +client = OpenAIModel("gpt-4.1-mini", base_url="https://gw.example/v1", api_key="sk-…") + +# Key from a named env var +client = OpenAIModel("gpt-4.1-mini", base_url="https://gw.example/v1", api_key_env="GW_KEY") + +# No args → uses OPENAI_API_KEY and the default OpenAI endpoint (backward-compatible) +client = OpenAIModel("gpt-4.1-mini") +``` + +Key-resolution precedence: explicit `api_key` → `api_key_env` env lookup → `OPENAI_API_KEY`. +Omitting `base_url` lets the OpenAI SDK default apply (honouring `OPENAI_BASE_URL` if set). + +> **Note:** `OpenAIModel` only calls `client.responses.create` (the Responses API). It is not +> suitable for Chat-Completions-only gateways. + +#### Ollama local and cloud usage + +`OllamaModel` supports both local Ollama runtimes and cloud endpoints. The `base_url` and +`api_key` / `api_key_env` keyword arguments give explicit control over the connection: + +```python +from prkit.core.model_clients.ollama import OllamaModel + +# Local (default: http://localhost:11434 or OLLAMA_HOST env) +client = OllamaModel("qwen3-vl:8b") + +# Local with explicit host +client = OllamaModel("qwen3-vl:8b", base_url="http://192.168.1.10:11434") + +# Cloud endpoint with explicit key +client = OllamaModel("llama3:70b-cloud", base_url="https://ollama.com", api_key="ol-…") + +# Cloud endpoint with key from env var +client = OllamaModel("llama3:70b-cloud", base_url="https://ollama.com", api_key_env="OLLAMA_CLOUD_KEY") + +# Env-var auth only (lib auto-reads OLLAMA_API_KEY when api_key/api_key_env not supplied) +client = OllamaModel("llama3:70b-cloud", base_url="https://ollama.com") +``` + +Key-resolution precedence: explicit `api_key` → `api_key_env` env lookup → library +auto-reads `OLLAMA_API_KEY`. For remote hosts (`base_url` pointing to a non-localhost +address) a failed startup preflight emits a warning instead of raising `ConnectionError`; +precise errors surface at `chat()` call time. + +#### Registering additional providers + +Use `register_model_client` to add new providers or override routing without modifying +built-in code: + +```python +from prkit.core.model_clients import register_model_client +from prkit.core.model_clients.factory import ProviderRule + +def _load_my_provider(model: str, logger): + from my_package import MyClient + return MyClient(model, logger) + +register_model_client(ProviderRule( + name="my_provider", + match=lambda model: model.startswith("my-"), + load=_load_my_provider, +)) +``` + ### PRKitLogger Centralized logger for consistent logging across PRKit packages. Provides colored console output, optional file logging, and environment-based configuration via `PRKIT_LOG_LEVEL`, `PRKIT_LOG_FILE`, `PRKIT_LOG_CONSOLE`, `PRKIT_LOG_COLORS`. Default log file: `{cwd}/prkit_logs/prkit.log`. diff --git a/DATASETS.md b/DATASETS.md index 163191f..ff9cdc3 100644 --- a/DATASETS.md +++ b/DATASETS.md @@ -545,3 +545,75 @@ To add a new dataset: 5. Add dataset information to this documentation See existing loaders in `src/prkit/datasets/loaders/` for examples. + +## Extending DatasetHub from External Code + +`DatasetHub` supports external loaders and downloaders registered at runtime — no fork or +subclass needed. + +### Registering an external loader + +```python +from prkit.datasets import DatasetHub +from prkit.datasets.loaders.base_loader import BaseDatasetLoader +from prkit.core.domain import PhysicalDataset, PhysicsProblem + +class MyLoader(BaseDatasetLoader): + @property + def field_mapping(self): + return {} + + def get_info(self): + return { + "name": "my_dataset", + "variants": ["full"], + "splits": ["full"], + } + + def load(self, data_dir=None, **kwargs): + # Read from data_dir and return a PhysicalDataset + ... + +DatasetHub.register("my_dataset", MyLoader) +``` + +After registration all hub methods (`load`, `get_info`, `list_available`) recognise +`"my_dataset"`. Built-in loaders are always present regardless of registration order. + +### Loading from a local directory (no downloader) + +A loader does **not** require a paired downloader. Pass `data_dir` to read from a local +path directly, bypassing any download step: + +```python +dataset = DatasetHub.load("my_dataset", data_dir="/path/to/data") +``` + +This works even when no `BaseDownloader` is registered for the name. + +### Registering an external downloader + +```python +from prkit.datasets.downloaders.base_downloader import BaseDownloader + +class MyDownloader(BaseDownloader): + @property + def dataset_name(self): + return "my_dataset" + + @property + def download_info(self): + return {"variants": ["full"], "splits": ["full"]} + + def _do_download(self, download_dir, **kwargs): + # Download logic — return download_dir when done + return download_dir + + def verify(self, data_dir): + return True + +DatasetHub.register_downloader("my_dataset", MyDownloader) +``` + +With a downloader registered, `DatasetHub.load("my_dataset", auto_download=True)` will +trigger `MyDownloader` when the data directory is missing. diff --git a/RELEASE_NOTES.md b/RELEASE_NOTES.md index 5e59d2d..968c4ee 100644 --- a/RELEASE_NOTES.md +++ b/RELEASE_NOTES.md @@ -1,3 +1,30 @@ +## Physical Reasoning Toolkit — Next Release + +### Highlights + +**Custom endpoint flexibility for model clients.** `OpenAIModel` now accepts `base_url`, +`api_key`, and `api_key_env` keyword arguments, making it straightforward to route traffic +to a proxy or gateway that fronts the OpenAI Responses API without subclassing. `OllamaModel` +gains the same `api_key` / `api_key_env` params for cloud endpoints (e.g. `ollama.com`), +and its startup connectivity check now treats remote hosts gracefully — a failed preflight +warns instead of raising, so cloud usage no longer requires suppressing the connection check. + +**`DatasetHub` registration-ordering bug fixed.** Calling `DatasetHub.register(name, Loader)` +before any built-in was touched previously caused all built-in loaders and downloaders to be +silently omitted. Built-ins are now seeded idempotently at the start of every public method. +External registrations can now happen in any order and are safe alongside built-in datasets. + +**Extending prkit — documented stable API.** `DATASETS.md` and `CORE.md` now document the +supported extension points: registering a `DatasetHub` loader or downloader from outside the +package, local-directory loading without a paired downloader, custom-endpoint construction for +`OpenAIModel` and `OllamaModel`, and adding new providers via `register_model_client`. + +**JSON-extraction consolidation (internal).** Three near-duplicate "extract JSON from model +text" implementations have been removed. All call sites delegate to the single tested +canonical helper in `prkit.core.model_clients.structured_output`. No public API change. + +--- + ## Physical Reasoning Toolkit v0.1.0 First release of **PRKit**—a unified toolkit for AI physical reasoning research. PRKit provides shared abstractions for physics problems, model inference, evaluation, and structured annotation workflows. diff --git a/pyproject.toml b/pyproject.toml index 4cac97f..e3ace38 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -1,5 +1,5 @@ [build-system] -requires = ["setuptools>=61.0", "wheel"] +requires = ["setuptools>=77.0", "wheel"] build-backend = "setuptools.build_meta" [project] @@ -7,7 +7,7 @@ name = "physical-reasoning-toolkit" version = "0.1.0.post22" description = "A toolkit for physical-reasoning datasets, multi-provider LLM inference, answer evaluation, and annotation." readme = {file = "README.md", content-type = "text/markdown"} -license = {text = "MIT"} +license = "MIT" authors = [ {name = "Yinghuan Zhang", email = "yinghuan.flash@gmail.com"} ] @@ -25,7 +25,6 @@ classifiers = [ "Intended Audience :: Developers", "Intended Audience :: Science/Research", "Intended Audience :: Education", - "License :: OSI Approved :: MIT License", "Operating System :: OS Independent", "Programming Language :: Python :: 3", "Programming Language :: Python :: 3.10", diff --git a/src/prkit/core/__init__.py b/src/prkit/core/__init__.py index 7cf78a1..3b0dafb 100644 --- a/src/prkit/core/__init__.py +++ b/src/prkit/core/__init__.py @@ -4,8 +4,20 @@ This package provides core functionality for PRKit (physical-reasoning-toolkit). """ +from .exceptions import ( + ConfigError, + DatasetError, + ModelClientError, + PRKitError, + UnknownModelError, +) from .logging_config import PRKitLogger __all__ = [ + "PRKitError", + "UnknownModelError", + "ModelClientError", + "ConfigError", + "DatasetError", "PRKitLogger", ] diff --git a/src/prkit/core/exceptions.py b/src/prkit/core/exceptions.py new file mode 100644 index 0000000..421a9e5 --- /dev/null +++ b/src/prkit/core/exceptions.py @@ -0,0 +1,31 @@ +""" +PRKit exception hierarchy. + +All package-specific exceptions derive from PRKitError so callers can catch +the entire family with a single ``except PRKitError`` clause while still +catching individual subtypes for finer-grained handling. + +Subclasses dual-inherit the closest matching builtin so that existing call +sites which assert on builtins (e.g. ``except ValueError``) continue to work +without modification. +""" + + +class PRKitError(Exception): + """Base class for all PRKit-specific exceptions.""" + + +class UnknownModelError(PRKitError, ValueError): + """Raised when a model name does not match any registered provider.""" + + +class ModelClientError(PRKitError, RuntimeError): + """Raised when a provider API call fails in a way that has useful context.""" + + +class ConfigError(PRKitError, ValueError): + """Raised for misconfigured or missing environment / config values.""" + + +class DatasetError(PRKitError, RuntimeError): + """Raised for dataset loading or download failures.""" diff --git a/src/prkit/core/model_clients/factory.py b/src/prkit/core/model_clients/factory.py index c219826..3589b74 100644 --- a/src/prkit/core/model_clients/factory.py +++ b/src/prkit/core/model_clients/factory.py @@ -16,6 +16,7 @@ from collections.abc import Callable from dataclasses import dataclass +from ..exceptions import UnknownModelError from .base import BaseModelClient ModelMatcher = Callable[[str], bool] @@ -76,7 +77,7 @@ def _load_gpt(model: str, logger: logging.Logger | None) -> BaseModelClient: from .openai import OpenAIModel, _is_supported_openai_model if not _is_supported_openai_model(model): - raise ValueError( + raise UnknownModelError( f"Unsupported OpenAI model: {model}. " "Supported OpenAI models: gpt-4.1, gpt-5xxxx (gpt-5.1, gpt-5.2, etc.), " "and o-family (o3, o4, o4-mini, etc.)" @@ -145,7 +146,7 @@ def create_model_client( for rule in _PROVIDER_RULES: if rule.matches(model_lower): return rule.load(model, logger) - raise ValueError( + raise UnknownModelError( f"Unknown model: {model}. " "Supported models: OpenAI (gpt-4.1, gpt-5xxxx, o-family), " "Anthropic (claude-*), Google (gemini-*), DeepSeek (deepseek-*), " diff --git a/src/prkit/core/model_clients/ollama.py b/src/prkit/core/model_clients/ollama.py index 8dab17f..ec9665f 100644 --- a/src/prkit/core/model_clients/ollama.py +++ b/src/prkit/core/model_clients/ollama.py @@ -12,6 +12,7 @@ import logging import os from typing import Any +from urllib.parse import urlparse import ollama @@ -24,6 +25,14 @@ ) +def _is_remote_host(url: str | None) -> bool: + """Return True when *url* points to a non-local host.""" + if not url: + return False + hostname = (urlparse(url).hostname or "").lower() + return hostname not in ("localhost", "127.0.0.1", "0.0.0.0") + + def normalize_ollama_model_name(model: str) -> str: """ Normalize Ollama model identifiers. @@ -68,6 +77,9 @@ def __init__( model: str, logger: logging.Logger | None = None, base_url: str | None = None, + *, + api_key: str | None = None, + api_key_env: str | None = None, ) -> None: """ Initialize Ollama model client. @@ -77,38 +89,72 @@ def __init__( - Raw model name: 'qwen3-vl', 'qwen2.5', 'llava' - Prefixed form: 'ollama/qwen3-vl:8b' logger: Optional logger instance - base_url: Optional base URL for Ollama API (defaults to http://localhost:11434) + base_url: Optional base URL for Ollama API (defaults to http://localhost:11434, + or ``OLLAMA_HOST`` env var). Use ``https://ollama.com`` for cloud. + api_key: Explicit API key forwarded as an ``Authorization: Bearer`` header. + Takes precedence over ``api_key_env``. When omitted the library falls + back to the ``OLLAMA_API_KEY`` environment variable automatically. + api_key_env: Name of an environment variable to read the API key from. + Used when ``api_key`` is not supplied. Raises: - ConnectionError: If Ollama service is not running or unreachable + ConnectionError: If Ollama service is not running or unreachable (local hosts only; + remote hosts emit a warning and continue). """ super().__init__(normalize_ollama_model_name(model), logger) self.provider = "ollama" self.base_url = base_url - # The ollama-python library uses a default client pointing to localhost:11434 - # but you can also use ollama.Client(host='...') if needed. + + # Resolve explicit auth — lib auto-injects OLLAMA_API_KEY when not set here. + if api_key is not None: + self._auth_header: str | None = f"Bearer {api_key}" + elif api_key_env is not None: + env_val = os.environ.get(api_key_env) + self._auth_header = f"Bearer {env_val}" if env_val else None + else: + self._auth_header = None # Check if Ollama is running during initialization self._check_ollama_running() + def _client_options(self) -> dict[str, Any]: + """Build keyword args for ``ollama.Client`` from resolved instance state.""" + opts: dict[str, Any] = {} + if self.base_url: + opts["host"] = self.base_url + if self._auth_header is not None: + # Use lowercase key so the lib's env-injection guard (which checks + # headers.get('authorization')) suppresses duplicate auth injection. + opts["headers"] = {"authorization": self._auth_header} + return opts + def _check_ollama_running(self) -> None: """ Check if Ollama service is running and accessible. + For remote hosts a failed preflight emits a warning instead of raising, + since network conditions differ and chat() will surface precise errors + at call time. + Raises: - ConnectionError: If Ollama service is not running or unreachable + ConnectionError: If Ollama service is not running on a local host. """ + effective_host = self.base_url or os.environ.get("OLLAMA_HOST") try: - # Try to list models as a simple connectivity check - if self.base_url: - client = ollama.Client(host=self.base_url) - else: - client = ollama.Client() + client = ollama.Client(**self._client_options()) client.list() except Exception as e: + if _is_remote_host(effective_host): + self.logger.warning( + "Remote Ollama host %s preflight failed (%s). " + "Continuing — errors will surface at inference time.", + effective_host, + e, + ) + return error_msg = ( f"Ollama service is not running or unreachable at " - f"{self.base_url or 'http://localhost:11434'}. " + f"{effective_host or 'http://localhost:11434'}. " f"Please ensure Ollama is installed and running.\n" f"To start Ollama:\n" f" 1. Install Ollama from https://ollama.com/download\n" @@ -195,10 +241,10 @@ def chat( if request_format is not None: request_kwargs["format"] = request_format - # Use Client if base_url is specified, otherwise use default - if self.base_url: - client = ollama.Client(host=self.base_url) - response = client.chat(**request_kwargs) + # Use Client when explicit options are set; fall back to module-level otherwise. + opts = self._client_options() + if opts: + response = ollama.Client(**opts).chat(**request_kwargs) else: response = ollama.chat(**request_kwargs) diff --git a/src/prkit/core/model_clients/openai.py b/src/prkit/core/model_clients/openai.py index 9bffcff..9dd51a1 100644 --- a/src/prkit/core/model_clients/openai.py +++ b/src/prkit/core/model_clients/openai.py @@ -10,6 +10,7 @@ """ import logging +import os from typing import Any from openai import OpenAI @@ -174,7 +175,15 @@ class OpenAIModel(BaseModelClient): supports_response_format_json_schema = True - def __init__(self, model: str, logger: logging.Logger | None = None) -> None: + def __init__( + self, + model: str, + logger: logging.Logger | None = None, + *, + base_url: str | None = None, + api_key: str | None = None, + api_key_env: str | None = None, + ) -> None: """ Initialize OpenAI model client. @@ -184,6 +193,13 @@ def __init__(self, model: str, logger: logging.Logger | None = None) -> None: - gpt-5xxxx (gpt-5, gpt-5.1, gpt-5.2, gpt-5.1-mini, etc.) - o-family (o3, o4, o4-mini, etc. - models starting with 'o' followed by number) logger: Optional logger instance + base_url: Optional custom Responses-API endpoint (e.g. a proxy at + ``https://gw.example/v1``). When omitted the OpenAI SDK default + is used (honouring its own ``OPENAI_BASE_URL`` env if set). + api_key: Explicit API key. Takes precedence over ``api_key_env`` and the + default ``OPENAI_API_KEY`` environment variable. + api_key_env: Name of an environment variable to read the API key from. + Used when ``api_key`` is not supplied. Raises: ValueError: If the model is not supported @@ -195,8 +211,21 @@ def __init__(self, model: str, logger: logging.Logger | None = None) -> None: "and o-family (o3, o4, o4-mini, etc.)" ) super().__init__(model, logger) - self.client = OpenAI(api_key=ensure_openai_api_key(__file__, required=False)) + + resolved_api_key: str | None + if api_key is not None: + resolved_api_key = api_key + elif api_key_env is not None: + resolved_api_key = os.environ.get(api_key_env) + else: + resolved_api_key = ensure_openai_api_key(__file__, required=False) + + client_kwargs: dict[str, Any] = {"api_key": resolved_api_key} + if base_url is not None: + client_kwargs["base_url"] = base_url + self.client = OpenAI(**client_kwargs) self.provider = "openai" + self.base_url = base_url self.is_o_family = _is_o_family_model(model) def chat( diff --git a/src/prkit/datasets/hub.py b/src/prkit/datasets/hub.py index ab7ace5..8b818d0 100644 --- a/src/prkit/datasets/hub.py +++ b/src/prkit/datasets/hub.py @@ -68,32 +68,38 @@ class DatasetHub: @classmethod def _register_default_loaders(cls) -> None: - """Register the default dataset loaders.""" - cls.register("physbench", PhysBenchLoader) - cls.register("phybench", PHYBenchLoader) - cls.register("physics", PhysicsLoader) - cls.register("phyx", PhyXLoader) - cls.register("seephys", SeePhysLoader) - cls.register("ugphysics", UGPhysicsLoader) - cls.register("jeebench", JEEBenchLoader) - cls.register("tpbench", TPBenchLoader) - cls.register("physreason", PhysReasonLoader) + """Register built-in loaders using setdefault so caller entries are never clobbered.""" + cls._loaders.setdefault("physbench", PhysBenchLoader) + cls._loaders.setdefault("phybench", PHYBenchLoader) + cls._loaders.setdefault("physics", PhysicsLoader) + cls._loaders.setdefault("phyx", PhyXLoader) + cls._loaders.setdefault("seephys", SeePhysLoader) + cls._loaders.setdefault("ugphysics", UGPhysicsLoader) + cls._loaders.setdefault("jeebench", JEEBenchLoader) + cls._loaders.setdefault("tpbench", TPBenchLoader) + cls._loaders.setdefault("physreason", PhysReasonLoader) @classmethod def _register_default_downloaders(cls) -> None: - """Register the default dataset downloaders.""" - cls.register_downloader("physbench", PhysBenchDownloader) - cls.register_downloader("phybench", PHYBenchDownloader) - cls.register_downloader("physics", PhysicsDownloader) - cls.register_downloader("phyx", PhyXDownloader) - cls.register_downloader("physreason", PhysReasonDownloader) - cls.register_downloader("seephys", SeePhysDownloader) - cls.register_downloader("ugphysics", UGPhysicsDownloader) - # Add more downloaders as they are implemented + """Register built-in downloaders using setdefault so caller entries are never clobbered.""" + cls._downloaders.setdefault("physbench", PhysBenchDownloader) + cls._downloaders.setdefault("phybench", PHYBenchDownloader) + cls._downloaders.setdefault("physics", PhysicsDownloader) + cls._downloaders.setdefault("phyx", PhyXDownloader) + cls._downloaders.setdefault("physreason", PhysReasonDownloader) + cls._downloaders.setdefault("seephys", SeePhysDownloader) + cls._downloaders.setdefault("ugphysics", UGPhysicsDownloader) + + @classmethod + def _ensure_defaults_registered(cls) -> None: + """Idempotently seed built-in loaders and downloaders.""" + cls._register_default_loaders() + cls._register_default_downloaders() @classmethod def register(cls, name: str, loader_class: type[BaseDatasetLoader]) -> None: """Register a new dataset loader.""" + cls._ensure_defaults_registered() cls._loaders[name] = loader_class @classmethod @@ -101,13 +107,13 @@ def register_downloader( cls, name: str, downloader_class: type[BaseDownloader] ) -> None: """Register a new dataset downloader.""" + cls._ensure_defaults_registered() cls._downloaders[name] = downloader_class @classmethod def _get_downloader(cls, name: str) -> BaseDownloader | None: """Get a dataset downloader by name.""" - if not cls._downloaders: - cls._register_default_downloaders() + cls._ensure_defaults_registered() if name not in cls._downloaders: return None @@ -154,8 +160,7 @@ def download( @classmethod def _get_loader(cls, name: str) -> BaseDatasetLoader: """Get a dataset loader by name.""" - if not cls._loaders: - cls._register_default_loaders() + cls._ensure_defaults_registered() if name not in cls._loaders: available = ", ".join(cls._loaders.keys()) @@ -376,8 +381,7 @@ def load( @classmethod def list_available(cls) -> list[str]: """List all available dataset names.""" - if not cls._loaders: - cls._register_default_loaders() + cls._ensure_defaults_registered() return list(cls._loaders.keys()) @classmethod diff --git a/src/prkit/datasets/loaders/base_loader.py b/src/prkit/datasets/loaders/base_loader.py index 1728eb6..95cf539 100644 --- a/src/prkit/datasets/loaders/base_loader.py +++ b/src/prkit/datasets/loaders/base_loader.py @@ -515,6 +515,33 @@ def _normalize_language(self, language: str) -> str: # Default fallback return "en" + @property + def DOMAIN_MAPPING(self) -> dict[str, Any]: + """Subclasses override to map raw domain strings to PhysicsDomain values.""" + return {} + + def _map_domain( + self, metadata: dict[str, Any], key: str = "domain" + ) -> dict[str, Any]: + """Normalize metadata[key] via DOMAIN_MAPPING, defaulting to PhysicsDomain.OTHER. + + Args: + metadata: Problem metadata dict (mutated in-place). + key: The metadata key holding the raw domain string. + + Returns: + The same metadata dict with metadata[key] replaced by a PhysicsDomain value, + or PhysicsDomain.OTHER when the key is absent or unmapped. + """ + from prkit.core.domain import PhysicsDomain + + raw = metadata.get(key) + if raw is not None: + metadata[key] = self.DOMAIN_MAPPING.get(raw, PhysicsDomain.OTHER) + else: + metadata[key] = PhysicsDomain.OTHER + return metadata + def validate_required_fields(self, data: dict[str, Any]) -> list[str]: """ Validate that problem data has required fields. diff --git a/src/prkit/datasets/loaders/phybench_loader.py b/src/prkit/datasets/loaders/phybench_loader.py index bd2212b..69d6f86 100644 --- a/src/prkit/datasets/loaders/phybench_loader.py +++ b/src/prkit/datasets/loaders/phybench_loader.py @@ -81,9 +81,7 @@ def _process_metadata(self, metadata: dict[str, Any]) -> dict[str, Any]: metadata["answer_category"] = "formula" - domain = metadata.get("domain") - if domain: - metadata["domain"] = self.DOMAIN_MAPPING.get(domain, PhysicsDomain.OTHER) + self._map_domain(metadata) return metadata diff --git a/src/prkit/datasets/loaders/phyx_loader.py b/src/prkit/datasets/loaders/phyx_loader.py index 22db56f..bd929f5 100644 --- a/src/prkit/datasets/loaders/phyx_loader.py +++ b/src/prkit/datasets/loaders/phyx_loader.py @@ -111,12 +111,7 @@ def _process_metadata(self, metadata: dict[str, Any]) -> dict[str, Any]: metadata["question"] = question # Map domain - domain = metadata.get("domain") - if domain: - normalized_domain = self.DOMAIN_MAPPING.get(domain, PhysicsDomain.OTHER) - metadata["domain"] = normalized_domain - else: - metadata["domain"] = PhysicsDomain.OTHER + self._map_domain(metadata) # Determine problem type options = metadata.get("options", []) diff --git a/src/prkit/datasets/loaders/tpbench_loader.py b/src/prkit/datasets/loaders/tpbench_loader.py index eefc35f..0ca66f8 100644 --- a/src/prkit/datasets/loaders/tpbench_loader.py +++ b/src/prkit/datasets/loaders/tpbench_loader.py @@ -88,9 +88,7 @@ def DOMAIN_MAPPING(self) -> dict[str, PhysicsDomain]: def _process_metadata(self, metadata: dict[str, Any]) -> dict[str, Any]: """Process metadata to create standardized problem fields.""" metadata["answer_category"] = "formula" - domain = metadata.get("domain") - if domain: - metadata["domain"] = self.DOMAIN_MAPPING.get(domain, PhysicsDomain.OTHER) + self._map_domain(metadata) return metadata diff --git a/src/prkit/evaluation/llm_judge/parse.py b/src/prkit/evaluation/llm_judge/parse.py index 189cd9a..5a182a2 100644 --- a/src/prkit/evaluation/llm_judge/parse.py +++ b/src/prkit/evaluation/llm_judge/parse.py @@ -2,10 +2,11 @@ from __future__ import annotations -import json -import re from typing import Any +from prkit.core.model_clients.structured_output import ( + extract_json_object as _canonical_extract_json_object, +) from prkit.evaluation.llm_judge.schema import EXPECTED_ANSWER_TYPES from prkit.evaluation.llm_judge.types import RESULT_SOURCE_LLM_JUDGE, LLMJudgeResult @@ -14,24 +15,7 @@ def extract_json_object(text: str) -> dict[str, Any] | None: """Return the first JSON object found in *text*, or ``None``.""" - raw = text.strip() - try: - parsed = json.loads(raw) - if isinstance(parsed, dict): - return parsed - except json.JSONDecodeError: - pass - - match = re.search(r"\{.*\}", raw, flags=re.DOTALL) - if not match: - return None - try: - parsed = json.loads(match.group(0)) - if isinstance(parsed, dict): - return parsed - except json.JSONDecodeError: - return None - return None + return _canonical_extract_json_object(text) def reasoning_indicates_incorrect(reasoning: str) -> bool: diff --git a/src/prkit/semantics/inference/calls.py b/src/prkit/semantics/inference/calls.py index a3b822b..74f3f20 100644 --- a/src/prkit/semantics/inference/calls.py +++ b/src/prkit/semantics/inference/calls.py @@ -2,10 +2,7 @@ from __future__ import annotations -import json import logging -import re -from collections.abc import Iterator from dataclasses import dataclass from pathlib import Path from typing import Any, TypeVar @@ -72,7 +69,6 @@ StrictReferenceSemanticsResponse, ) -_JSON_FENCE_RE = re.compile(r"```(?:json)?\s*(.*?)\s*```", re.DOTALL) logger = logging.getLogger(__name__) _STRICT_PREDICTION_RESPONSE_FIELDS = frozenset( StrictPredictionSemanticsResponse.model_fields @@ -585,7 +581,7 @@ def _parse_response_model( try: return response_model.model_validate_json(text) except ValidationError: - payload = _extract_json_object(text) + payload = extract_structured_json_object(text) if payload is None: raise ValueError( f"Could not parse {response_model.__name__} response as JSON.\n" @@ -700,58 +696,6 @@ def _normalize_response_payload( return normalized -def _extract_json_object(text: str) -> dict[str, Any] | None: - """Extract the first JSON object from a model response, if any.""" - return extract_structured_json_object(text) - - -def _iter_braced_json_candidates(text: str) -> Iterator[str]: - """Yield balanced brace substrings that could be JSON objects.""" - - start_indices = [index for index, char in enumerate(text) if char == "{"] - for start in start_indices: - depth = 0 - in_string = False - escape = False - for index in range(start, len(text)): - char = text[index] - if in_string: - if escape: - escape = False - elif char == "\\": - escape = True - elif char == '"': - in_string = False - continue - if char == '"': - in_string = True - continue - if char == "{": - depth += 1 - elif char == "}": - depth -= 1 - if depth == 0: - yield text[start : index + 1] - break - - -def _try_parse_json_object(candidate: str) -> dict[str, Any] | None: - """Parse one JSON-object candidate and ignore nested-shape false positives.""" - - try: - parsed = json.loads(candidate) - except json.JSONDecodeError: - return None - if not isinstance(parsed, dict): - return None - if _looks_like_prediction_answer_semantics_payload(parsed): - return { - "prediction_answer_semantics": parsed, - "question_semantics": {}, - } - return parsed - - def _looks_like_prediction_answer_semantics_payload(payload: dict[str, Any]) -> bool: """Whether a parsed object looks like only the nested answer-semantics block.""" diff --git a/technical_reports/evaluation_subset_selection_memo_20260414.md b/technical_reports/evaluation_subset_selection_memo_20260414.md index dd86965..6fd5e33 100644 --- a/technical_reports/evaluation_subset_selection_memo_20260414.md +++ b/technical_reports/evaluation_subset_selection_memo_20260414.md @@ -62,11 +62,11 @@ Avoid claiming any of the following unless a provenance script or note is recove ## Repo anchors - Official split metadata lives in: - - `src/prkit/prkit_datasets/loaders/seephys_loader.py` - - `src/prkit/prkit_datasets/loaders/phyx_loader.py` - - `src/prkit/prkit_datasets/loaders/physreason_loader.py` - - `src/prkit/prkit_datasets/loaders/physbench_loader.py` - - `src/prkit/prkit_datasets/loaders/physics_loader.py` - - `src/prkit/prkit_datasets/loaders/ugphysics_loader.py` + - `src/prkit/datasets/loaders/seephys_loader.py` + - `src/prkit/datasets/loaders/phyx_loader.py` + - `src/prkit/datasets/loaders/physreason_loader.py` + - `src/prkit/datasets/loaders/physbench_loader.py` + - `src/prkit/datasets/loaders/physics_loader.py` + - `src/prkit/datasets/loaders/ugphysics_loader.py` - Fixed evaluation ID lists live in: - `uncertainty_quantification_physical_reasoning/perturbations//problem_ids_for_perturbation.json` diff --git a/technical_reports/predicate_comparison_technical_report.md b/technical_reports/predicate_comparison_technical_report.md index c5f7bf5..d90d0d6 100644 --- a/technical_reports/predicate_comparison_technical_report.md +++ b/technical_reports/predicate_comparison_technical_report.md @@ -6,7 +6,7 @@ This report uses the name `predicate` throughout. In the codebase, the predicate comparator corresponds to: -- module: `src/prkit/prkit_evaluation/comparator/smart_llm.py` +- module: `src/prkit/evaluation/comparator/smart_llm.py` - class: `SmartLLMComparator` - builder name: `build_comparator("smart_llm")` @@ -30,7 +30,7 @@ This is stricter and more physics-aware than plain string matching, but cheaper ## Package Map -The predicate path sits inside the broader `prkit_evaluation` package: +The predicate path sits inside the broader `prkit.evaluation` package: - `comparator/` - comparator abstractions and concrete answer comparators @@ -45,20 +45,20 @@ The predicate path sits inside the broader `prkit_evaluation` package: Important files for the predicate path: -- `src/prkit/prkit_evaluation/comparator/base.py` -- `src/prkit/prkit_evaluation/comparator/by_module.py` -- `src/prkit/prkit_evaluation/comparator/smart_llm.py` -- `src/prkit/prkit_evaluation/comparator/smart_match.py` -- `src/prkit/prkit_evaluation/comparator/smart_pipeline.py` -- `src/prkit/prkit_evaluation/comparator/typed_llm.py` -- `src/prkit/prkit_evaluation/utils/normalization.py` -- `src/prkit/prkit_evaluation/utils/compare_same_type.py` -- `src/prkit/prkit_evaluation/utils/compare_cross_type.py` -- `src/prkit/prkit_evaluation/utils/category_dispatch.py` -- `src/prkit/prkit_evaluation/utils/answer_utils.py` -- `src/prkit/prkit_evaluation/llm_judge/*` -- `src/prkit/prkit_evaluation/evaluator/base.py` -- `src/prkit/prkit_evaluation/evaluator/accuracy.py` +- `src/prkit/evaluation/comparator/base.py` +- `src/prkit/evaluation/comparator/by_module.py` +- `src/prkit/evaluation/comparator/smart_llm.py` +- `src/prkit/evaluation/comparator/smart_match.py` +- `src/prkit/evaluation/comparator/smart_pipeline.py` +- `src/prkit/evaluation/comparator/typed_llm.py` +- `src/prkit/evaluation/utils/normalization.py` +- `src/prkit/evaluation/utils/compare_same_type.py` +- `src/prkit/evaluation/utils/compare_cross_type.py` +- `src/prkit/evaluation/utils/category_dispatch.py` +- `src/prkit/evaluation/utils/answer_utils.py` +- `src/prkit/evaluation/llm_judge/*` +- `src/prkit/evaluation/evaluator/base.py` +- `src/prkit/evaluation/evaluator/accuracy.py` ## Core Interface @@ -84,7 +84,7 @@ This orientation matters in cross-type rules, especially around unit handling an The standard factory is: ```python -from prkit.prkit_evaluation.comparator import build_comparator +from prkit.evaluation.comparator import build_comparator predicate = build_comparator("smart_llm") ``` @@ -96,14 +96,14 @@ Relevant facts: - `smart_llm` maps to `SmartLLMComparator` - `typed_llm` maps to `TypedLLMComparator` - both accept an OpenAI model name -- default judge model comes from `prkit.prkit_evaluation.llm_judge.DEFAULT_MODEL` +- default judge model comes from `prkit.evaluation.llm_judge.DEFAULT_MODEL` ## Input Representation The comparator accepts either: - raw strings -- `Answer` objects from `prkit.prkit_core.domain.answer` +- `Answer` objects from `prkit.core.domain.answer` If an `Answer` object is provided, its existing category and value are reused. If a raw string is provided, the comparator normalizes and categorizes it at runtime. @@ -485,7 +485,7 @@ The broader package includes several comparator families: - predicate comparator - SmartMatch pipeline plus LLM fallback only on true deterministic inconclusiveness -This explains the role of the other `prkit_evaluation` modules relative to predicate comparison. +This explains the role of the other `prkit.evaluation` modules relative to predicate comparison. ## Evaluator Integration @@ -501,8 +501,8 @@ The predicate comparator is consumed through `evaluator/accuracy.py`. For predicate comparison, typical use is: ```python -from prkit.prkit_evaluation.comparator import build_comparator -from prkit.prkit_evaluation.evaluator import AccuracyEvaluator +from prkit.evaluation.comparator import build_comparator +from prkit.evaluation.evaluator import AccuracyEvaluator predicate = build_comparator("smart_llm") evaluator = AccuracyEvaluator(predicate) @@ -514,7 +514,7 @@ result = evaluator.evaluate(predicted_answer, ground_truth_answer, question=ques ### Deterministic Predicate-Only Path ```python -from prkit.prkit_evaluation.comparator import build_comparator +from prkit.evaluation.comparator import build_comparator predicate = build_comparator("smart_llm") score = predicate.accuracy_score( @@ -534,7 +534,7 @@ Behavior: ### Full Predicate Path With Judge Fallback ```python -from prkit.prkit_evaluation.comparator import build_comparator +from prkit.evaluation.comparator import build_comparator predicate = build_comparator("smart_llm", model="gpt-5.4-mini") matched = predicate.compare( diff --git a/tests/prkit/core/model_clients/test_factory.py b/tests/prkit/core/model_clients/test_factory.py index a13d361..23578ce 100644 --- a/tests/prkit/core/model_clients/test_factory.py +++ b/tests/prkit/core/model_clients/test_factory.py @@ -35,6 +35,16 @@ class TestCreateModelClient: """Test cases for create_model_client factory function.""" + @pytest.fixture(autouse=True) + def _mock_provider_sdks(self): + """Prevent real SDK credential checks during factory routing tests.""" + with ( + patch("prkit.core.model_clients.openai.OpenAI"), + patch("prkit.core.model_clients.gemini.genai"), + patch("prkit.core.model_clients.openai_compatible_chat.OpenAI"), + ): + yield + def test_create_openai_gpt_4_1(self): """Test creating OpenAI gpt-4.1 model.""" client = create_model_client("gpt-4.1") diff --git a/tests/prkit/core/model_clients/test_ollama.py b/tests/prkit/core/model_clients/test_ollama.py index 29a20c5..1df3534 100644 --- a/tests/prkit/core/model_clients/test_ollama.py +++ b/tests/prkit/core/model_clients/test_ollama.py @@ -2,6 +2,7 @@ Tests for Ollama model client. """ +import os from unittest.mock import MagicMock, Mock, patch import pytest @@ -352,3 +353,83 @@ def test_chat_with_custom_logger(self, mock_ollama_module): client = OllamaModel(OLLAMA_QWEN_TEST_MODEL, logger=logger) assert client.logger == logger + + +class TestOllamaModelCloudAuth: + """Tests for explicit api_key / api_key_env params and remote-safe preflight.""" + + @patch("prkit.core.model_clients.ollama.ollama") + def test_explicit_api_key_forwarded_as_auth_header(self, mock_ollama_module): + """Explicit api_key is sent as a lowercase authorization header to ollama.Client.""" + mock_client = MagicMock() + mock_client.list.return_value = [] + mock_ollama_module.Client.return_value = mock_client + + OllamaModel( + OLLAMA_QWEN_TEST_MODEL, + base_url="https://ollama.com", + api_key="test-cloud-key", + ) + + _, kwargs = mock_ollama_module.Client.call_args + assert kwargs["host"] == "https://ollama.com" + assert kwargs["headers"] == {"authorization": "Bearer test-cloud-key"} + + @patch("prkit.core.model_clients.ollama.ollama") + def test_api_key_env_resolves_named_var(self, mock_ollama_module): + """api_key_env reads the key from the named environment variable.""" + mock_client = MagicMock() + mock_client.list.return_value = [] + mock_ollama_module.Client.return_value = mock_client + + with patch.dict(os.environ, {"MY_OLLAMA_KEY": "env-ollama-key"}): + OllamaModel( + OLLAMA_QWEN_TEST_MODEL, + base_url="https://ollama.com", + api_key_env="MY_OLLAMA_KEY", + ) + + _, kwargs = mock_ollama_module.Client.call_args + assert kwargs["headers"] == {"authorization": "Bearer env-ollama-key"} + + @patch("prkit.core.model_clients.ollama.ollama") + def test_no_api_key_falls_back_to_module_level_chat(self, mock_ollama_module): + """Without base_url or api_key, chat falls back to module-level ollama.chat.""" + mock_client = MagicMock() + mock_client.list.return_value = [] + mock_ollama_module.Client.return_value = mock_client + + mock_response = Mock() + mock_response.message = Mock() + mock_response.message.content = "Response" + mock_ollama_module.chat.return_value = mock_response + + client = OllamaModel(OLLAMA_QWEN_TEST_MODEL) + client.chat("Hello") + + mock_ollama_module.chat.assert_called_once() + + @patch("prkit.core.model_clients.ollama.ollama") + def test_remote_host_preflight_failure_warns_not_raises(self, mock_ollama_module): + """A failed preflight against a remote host emits a warning but does not raise.""" + mock_client = MagicMock() + mock_client.list.side_effect = Exception("Network unreachable") + mock_ollama_module.Client.return_value = mock_client + + # Should not raise — remote host gets a warning instead + client = OllamaModel( + OLLAMA_QWEN_TEST_MODEL, + base_url="https://ollama.com", + api_key="k", + ) + assert client.base_url == "https://ollama.com" + + @patch("prkit.core.model_clients.ollama.ollama") + def test_local_host_preflight_failure_still_raises(self, mock_ollama_module): + """A failed preflight against a local host still raises ConnectionError.""" + mock_client = MagicMock() + mock_client.list.side_effect = Exception("Connection refused") + mock_ollama_module.Client.return_value = mock_client + + with pytest.raises(ConnectionError, match="Ollama service is not running"): + OllamaModel(OLLAMA_QWEN_TEST_MODEL) diff --git a/tests/prkit/core/model_clients/test_openai.py b/tests/prkit/core/model_clients/test_openai.py index 4dd9c69..2bf033e 100644 --- a/tests/prkit/core/model_clients/test_openai.py +++ b/tests/prkit/core/model_clients/test_openai.py @@ -2,6 +2,7 @@ Tests for OpenAI model client. """ +import os from unittest.mock import MagicMock, Mock, patch import pytest @@ -66,15 +67,19 @@ def test_is_o_family_model_false(self): class TestOpenAIModel: """Test cases for OpenAIModel class.""" - def test_init_supported_model(self): + @patch("prkit.core.model_clients.openai.OpenAI") + def test_init_supported_model(self, mock_openai_class): """Test initializing with supported model.""" + mock_openai_class.return_value = MagicMock() client = OpenAIModel(OPENAI_TEST_MODEL) assert client.model == OPENAI_TEST_MODEL assert client.provider == "openai" assert client.is_o_family is False - def test_init_o_family_model(self): + @patch("prkit.core.model_clients.openai.OpenAI") + def test_init_o_family_model(self, mock_openai_class): """Test initializing with o-family model.""" + mock_openai_class.return_value = MagicMock() client = OpenAIModel("o3") assert client.model == "o3" assert client.is_o_family is True @@ -421,3 +426,37 @@ def test_is_o_family_model_edge_cases(self): assert _is_o_family_model("openai") is False # 'o' but not followed by digit assert _is_o_family_model("o") is False # Too short assert _is_o_family_model("oa") is False # 'o' followed by letter + + +class TestOpenAIModelCustomEndpoint: + """Test custom endpoint / API key params on OpenAIModel.""" + + @patch("prkit.core.model_clients.openai.OpenAI") + def test_explicit_base_url_and_api_key_forwarded(self, mock_openai_class): + """Explicit base_url and api_key are passed directly to the OpenAI SDK client.""" + mock_openai_class.return_value = MagicMock() + OpenAIModel( + OPENAI_TEST_MODEL, + base_url="https://gw.example/v1", + api_key="explicit-key", + ) + _, kwargs = mock_openai_class.call_args + assert kwargs["base_url"] == "https://gw.example/v1" + assert kwargs["api_key"] == "explicit-key" + + @patch("prkit.core.model_clients.openai.OpenAI") + def test_api_key_env_resolves_named_var(self, mock_openai_class): + """api_key_env reads the key from the named environment variable.""" + mock_openai_class.return_value = MagicMock() + with patch.dict(os.environ, {"MY_CUSTOM_KEY": "env-key-value"}): + OpenAIModel(OPENAI_TEST_MODEL, api_key_env="MY_CUSTOM_KEY") + _, kwargs = mock_openai_class.call_args + assert kwargs["api_key"] == "env-key-value" + + @patch("prkit.core.model_clients.openai.OpenAI") + def test_omitting_base_url_does_not_forward_it(self, mock_openai_class): + """Omitting base_url must not pass base_url= to the SDK (backward-compat guard).""" + mock_openai_class.return_value = MagicMock() + OpenAIModel(OPENAI_TEST_MODEL) + _, kwargs = mock_openai_class.call_args + assert "base_url" not in kwargs diff --git a/tests/prkit/core/test_exceptions.py b/tests/prkit/core/test_exceptions.py new file mode 100644 index 0000000..c8c762b --- /dev/null +++ b/tests/prkit/core/test_exceptions.py @@ -0,0 +1,92 @@ +"""Tests for prkit.core.exceptions hierarchy and factory wiring.""" + +from __future__ import annotations + +import pytest + +from prkit.core import ( + ModelClientError as CoreModelClientError, +) +from prkit.core import ( + PRKitError as CorePRKitError, +) +from prkit.core import ( + UnknownModelError as CoreUnknownModelError, +) +from prkit.core.exceptions import ( + ConfigError, + DatasetError, + ModelClientError, + PRKitError, + UnknownModelError, +) + + +class TestInheritanceContract: + def test_unknown_model_error_is_prkit_error(self): + exc = UnknownModelError("bad-model") + assert isinstance(exc, PRKitError) + + def test_unknown_model_error_is_value_error(self): + exc = UnknownModelError("bad-model") + assert isinstance(exc, ValueError) + + def test_model_client_error_is_prkit_error(self): + exc = ModelClientError("failed") + assert isinstance(exc, PRKitError) + + def test_model_client_error_is_runtime_error(self): + exc = ModelClientError("failed") + assert isinstance(exc, RuntimeError) + + def test_config_error_is_prkit_error(self): + assert isinstance(ConfigError("bad config"), PRKitError) + + def test_config_error_is_value_error(self): + assert isinstance(ConfigError("bad config"), ValueError) + + def test_dataset_error_is_prkit_error(self): + assert isinstance(DatasetError("load failed"), PRKitError) + + def test_dataset_error_is_runtime_error(self): + assert isinstance(DatasetError("load failed"), RuntimeError) + + def test_all_subtypes_catchable_as_prkit_error(self): + for exc_cls in (UnknownModelError, ModelClientError, ConfigError, DatasetError): + with pytest.raises(PRKitError): + raise exc_cls("test") + + +class TestCoreReExports: + """Ensure the public surface in prkit.core re-exports the same objects.""" + + def test_prkit_error_re_exported(self): + assert CorePRKitError is PRKitError + + def test_unknown_model_error_re_exported(self): + assert CoreUnknownModelError is UnknownModelError + + def test_model_client_error_re_exported(self): + assert CoreModelClientError is ModelClientError + + +class TestFactoryRaisesUnknownModelError: + """factory.create_model_client must raise UnknownModelError for bad model names.""" + + def test_completely_unknown_model(self): + from prkit.core.model_clients import create_model_client + + with pytest.raises(UnknownModelError, match="Unknown model"): + create_model_client("no-such-provider-xyz") + + def test_unknown_model_also_caught_as_value_error(self): + from prkit.core.model_clients import create_model_client + + with pytest.raises(ValueError): + create_model_client("no-such-provider-xyz") + + def test_unsupported_gpt_variant(self): + from prkit.core.model_clients import create_model_client + + with pytest.raises(UnknownModelError, match="Unsupported OpenAI model"): + create_model_client("gpt-3-turbo") diff --git a/tests/prkit/datasets/loaders/test_base_loader_map_domain.py b/tests/prkit/datasets/loaders/test_base_loader_map_domain.py new file mode 100644 index 0000000..cdfed60 --- /dev/null +++ b/tests/prkit/datasets/loaders/test_base_loader_map_domain.py @@ -0,0 +1,119 @@ +"""Tests for BaseDatasetLoader._map_domain helper (B4).""" + +from __future__ import annotations + +from prkit.core.domain import PhysicalDataset, PhysicsDomain +from prkit.datasets.loaders.base_loader import BaseDatasetLoader + + +class _LoaderNoDomainMapping(BaseDatasetLoader): + """Loader with no DOMAIN_MAPPING (inherits empty dict from base).""" + + @property + def field_mapping(self) -> dict[str, str]: + return {} + + def load(self, data_dir, **kwargs) -> PhysicalDataset: # type: ignore[override] + return PhysicalDataset(problems=[]) + + def get_info(self) -> dict: + return {} + + +class _LoaderWithDomainMapping(BaseDatasetLoader): + """Loader with a concrete DOMAIN_MAPPING property.""" + + @property + def DOMAIN_MAPPING(self) -> dict[str, PhysicsDomain]: + return { + "Mechanics": PhysicsDomain.MECHANICS, + "Thermodynamics": PhysicsDomain.THERMODYNAMICS, + } + + @property + def field_mapping(self) -> dict[str, str]: + return {} + + def load(self, data_dir, **kwargs) -> PhysicalDataset: # type: ignore[override] + return PhysicalDataset(problems=[]) + + def get_info(self) -> dict: + return {} + + +class TestMapDomainBase: + def setup_method(self): + self.loader = _LoaderNoDomainMapping() + + def test_absent_domain_sets_other(self): + meta: dict = {} + result = self.loader._map_domain(meta) + assert result["domain"] is PhysicsDomain.OTHER + + def test_unknown_domain_sets_other(self): + meta = {"domain": "UnknownSubfield"} + result = self.loader._map_domain(meta) + assert result["domain"] is PhysicsDomain.OTHER + + def test_returns_same_dict(self): + meta: dict = {} + result = self.loader._map_domain(meta) + assert result is meta + + +class TestMapDomainWithMapping: + def setup_method(self): + self.loader = _LoaderWithDomainMapping() + + def test_known_domain_maps_correctly(self): + meta = {"domain": "Mechanics"} + self.loader._map_domain(meta) + assert meta["domain"] is PhysicsDomain.MECHANICS + + def test_second_known_domain(self): + meta = {"domain": "Thermodynamics"} + self.loader._map_domain(meta) + assert meta["domain"] is PhysicsDomain.THERMODYNAMICS + + def test_unmapped_domain_falls_back_to_other(self): + meta = {"domain": "Biology"} + self.loader._map_domain(meta) + assert meta["domain"] is PhysicsDomain.OTHER + + def test_absent_domain_sets_other(self): + meta: dict = {} + self.loader._map_domain(meta) + assert meta["domain"] is PhysicsDomain.OTHER + + def test_custom_key(self): + meta = {"subject": "Mechanics"} + self.loader._map_domain(meta, key="subject") + assert meta["subject"] is PhysicsDomain.MECHANICS + + +class TestLoaderOverridePreserved: + """Verify the 3 concrete loaders still behave correctly after refactor.""" + + def test_phyx_loader_maps_domain(self): + from prkit.datasets.loaders.phyx_loader import PhyXLoader + + loader = PhyXLoader() + meta = {"domain": "Mechanics"} + loader._map_domain(meta) + assert meta["domain"] is PhysicsDomain.MECHANICS + + def test_phybench_loader_maps_domain(self): + from prkit.datasets.loaders.phybench_loader import PHYBenchLoader + + loader = PHYBenchLoader() + meta = {"domain": "MECHANICS"} + loader._map_domain(meta) + assert meta["domain"] is PhysicsDomain.MECHANICS + + def test_tpbench_loader_maps_domain(self): + from prkit.datasets.loaders.tpbench_loader import TPBenchLoader + + loader = TPBenchLoader() + meta = {"domain": "QM"} + loader._map_domain(meta) + assert meta["domain"] is PhysicsDomain.QUANTUM_MECHANICS diff --git a/tests/prkit/datasets/test_hub.py b/tests/prkit/datasets/test_hub.py index cb19270..0983cd3 100644 --- a/tests/prkit/datasets/test_hub.py +++ b/tests/prkit/datasets/test_hub.py @@ -736,6 +736,96 @@ def test_get_downloader_nonexistent(self): assert downloader is None +class TestDatasetHubExtensionContract: + """Tests for the external-loader extension contract (Task 3b).""" + + def _make_dummy_loader(self, tmp_path): + """Return a minimal BaseDatasetLoader that reads from *tmp_path*.""" + + class DummyExternalLoader(BaseDatasetLoader): + @property + def field_mapping(self): + return {} + + def get_info(self): + return { + "name": "dummy_ext", + "variants": ["full"], + "splits": ["full"], + } + + def load(self, data_dir=None, **kwargs): + import json + from pathlib import Path + + items = json.loads((Path(data_dir) / "problems.json").read_text()) + return PhysicalDataset( + problems=[ + PhysicsProblem(problem_id=item["id"], question=item["q"]) + for item in items + ] + ) + + return DummyExternalLoader + + def test_external_loader_loads_from_local_dir(self, tmp_path): + """External loader registered via DatasetHub.register can load from a local directory.""" + import json + + (tmp_path / "problems.json").write_text( + json.dumps( + [{"id": "p1", "q": "What is gravity?"}, {"id": "p2", "q": "F=ma?"}] + ) + ) + + DummyExternalLoader = self._make_dummy_loader(tmp_path) + DatasetHub.register("dummy_ext", DummyExternalLoader) + + try: + dataset = DatasetHub.load("dummy_ext", data_dir=tmp_path) + assert len(dataset) == 2 + assert dataset[0].problem_id == "p1" + assert dataset[1].question == "F=ma?" + finally: + DatasetHub._loaders.pop("dummy_ext", None) + + def test_external_registration_preserves_builtins(self): + """Registering an external loader before any built-in is touched still exposes built-ins.""" + saved_loaders = dict(DatasetHub._loaders) + saved_downloaders = dict(DatasetHub._downloaders) + + DatasetHub._loaders.clear() + DatasetHub._downloaders.clear() + + class DummyExternalLoader(BaseDatasetLoader): + @property + def field_mapping(self): + return {} + + def get_info(self): + return {"name": "dummy_guard", "variants": ["full"], "splits": ["full"]} + + def load(self, data_dir=None, **kwargs): + return PhysicalDataset(problems=[]) + + try: + # External loader registered first, with _loaders empty + DatasetHub.register("dummy_guard", DummyExternalLoader) + + available = DatasetHub.list_available() + assert ( + "ugphysics" in available + ), "Built-in loaders must be present after external registration" + assert ( + "dummy_guard" in available + ), "External loader must survive built-in registration" + finally: + DatasetHub._loaders.clear() + DatasetHub._loaders.update(saved_loaders) + DatasetHub._downloaders.clear() + DatasetHub._downloaders.update(saved_downloaders) + + class TestDatasetLoadersIntegration: """Integration tests for dataset loaders.""" diff --git a/tests/prkit/evaluation/comparator/test_by_module.py b/tests/prkit/evaluation/comparator/test_by_module.py index 276e6d7..1ab9da1 100644 --- a/tests/prkit/evaluation/comparator/test_by_module.py +++ b/tests/prkit/evaluation/comparator/test_by_module.py @@ -2,6 +2,8 @@ from __future__ import annotations +from unittest.mock import patch + import pytest from prkit.evaluation.comparator.by_module import ( @@ -31,7 +33,8 @@ def test_build_exact_match() -> None: assert isinstance(c, ExactMatchComparator) -def test_build_typed_llm_uses_model() -> None: +@patch("prkit.evaluation.llm_judge.runner.OpenAI") +def test_build_typed_llm_uses_model(mock_openai) -> None: c = build_comparator("typed_llm", model="gpt-4.1-mini") assert isinstance(c, TypedLLMComparator) assert c.model_name == "gpt-4.1-mini" diff --git a/tests/prkit/evaluation/comparator/test_smart_llm.py b/tests/prkit/evaluation/comparator/test_smart_llm.py index c63a680..0d0cb06 100644 --- a/tests/prkit/evaluation/comparator/test_smart_llm.py +++ b/tests/prkit/evaluation/comparator/test_smart_llm.py @@ -1,5 +1,7 @@ """Tests for SmartLLMComparator: deterministic SmartMatch path + LLM fallback metadata.""" +from unittest.mock import patch + import pytest from prkit.evaluation.comparator.smart_llm import SmartLLMComparator @@ -10,6 +12,11 @@ class TestSmartLLMComparator: + @pytest.fixture(autouse=True) + def _mock_openai(self): + with patch("prkit.evaluation.llm_judge.runner.OpenAI"): + yield + def test_match_records_smart_match_source(self): comp = SmartLLMComparator() assert comp.compare("42", "42") is True diff --git a/tests/uq/conftest.py b/tests/uq/conftest.py new file mode 100644 index 0000000..77891ea --- /dev/null +++ b/tests/uq/conftest.py @@ -0,0 +1,11 @@ +import importlib.util + +_uq_available = importlib.util.find_spec("uncertainty_quantification_physical_reasoning") is not None + +collect_ignore: list[str] = [] +if not _uq_available: + collect_ignore = [ + "test_batch_prepare_common.py", + "test_extract_answer_parsers.py", + "test_single_inference_with_answer_tags.py", + ]