chore: merge executor worktree (worktree-agent-afe278fe2f809d30b)

This commit is contained in:
curo1305
2026-06-04 19:04:52 +02:00
10 changed files with 706 additions and 80 deletions
@@ -0,0 +1,182 @@
---
phase: 07-redo-and-optimize-llm-integration
plan: "02"
subsystem: backend/ai-providers
tags:
- ai
- provider-refactor
- singleton-client
- json-mode
- smart-truncation
- registry-pattern
- pydantic
- wave-2
dependency_graph:
requires:
- "07-01 (system_settings table, HKDF helpers, xfail stubs)"
provides:
- "ProviderConfig Pydantic model + PROVIDER_DEFAULTS + SUPPORTS_JSON_MODE"
- "GenericOpenAIProvider covering all 8 OpenAI-compat vendors"
- "OpenAIProvider singleton _client lifecycle (D-07)"
- "Smart truncation _truncate() 60/40 (D-13)"
- "Registry-based get_provider(config: ProviderConfig) — no if/elif chain"
- "MAX_AI_CHARS removed from openai_provider.py and classifier.py"
- "anthropic SDK floor bumped to >=0.95.0"
- "6 Wave-2 xfail stubs promoted to passing tests"
affects:
- "07-03 (Anthropic output_config — depends on ProviderConfig and registry)"
- "07-04 (Celery retry — depends on classifier.py clean pass-through)"
- "07-05 (Admin AI panel — depends on ProviderConfig for form validation)"
tech_stack:
added: []
patterns:
- "ProviderConfig(BaseModel) with extra=forbid; context_chars=0 sentinel for PROVIDER_DEFAULTS resolution"
- "PROVIDER_DEFAULTS dict with 10 entries; SUPPORTS_JSON_MODE dict with gemini=False"
- "GenericOpenAIProvider(OpenAIProvider) with conditional response_format kwarg (D-01/D-02)"
- "Singleton self._client = AsyncOpenAI(...) in __init__ (D-07)"
- "_truncate(): first 60% + last 40% of context_chars (D-13)"
- "_REGISTRY dict in ai/__init__.py replaces if/elif chain (O(1) lookup)"
- "D-02 invariant: parse_classification/parse_suggestions always imported from ai.utils"
key_files:
created:
- backend/ai/provider_config.py
- backend/ai/generic_openai_provider.py
modified:
- backend/ai/openai_provider.py
- backend/ai/ollama_provider.py
- backend/ai/lmstudio_provider.py
- backend/ai/__init__.py
- backend/services/classifier.py
- backend/requirements.txt
- backend/tests/test_ai_providers.py
decisions:
- "context_chars=0 as sentinel in ProviderConfig (not 8000) — enables factory to resolve PROVIDER_DEFAULTS via `config.context_chars or defaults['context_chars']`"
- "GenericOpenAIProvider always calls parse_classification() as last-resort regardless of json_mode — D-02 invariant preserved even when json_object is requested"
- "ai/__init__.py imports GenericOpenAIProvider at module load time; no lazy import needed since provider_config.py has no side effects"
- "anthropic floor bumped to >=0.95.0 to unblock Plan 03 output_config usage (A5 from RESEARCH.md)"
metrics:
duration: "~35 minutes"
completed: "2026-06-04"
tasks_completed: 4
tasks_total: 4
files_created: 2
files_modified: 7
---
# Phase 7 Plan 02: Provider Config, GenericOpenAIProvider, and Registry Factory Summary
ProviderConfig Pydantic model with 10-provider PROVIDER_DEFAULTS + SUPPORTS_JSON_MODE; GenericOpenAIProvider subclassing OpenAIProvider with JSON-mode conditional on supports_json_mode; singleton _client lifecycle; smart truncation; registry-based get_provider(); MAX_AI_CHARS removed from two files; 6 Wave-2 xfail tests promoted.
## Tasks Completed
| Task | Description | Commit | Files |
|------|-------------|--------|-------|
| 1 | ProviderConfig + PROVIDER_DEFAULTS + SUPPORTS_JSON_MODE | beb5b5e | backend/ai/provider_config.py |
| 2 | GenericOpenAIProvider + singleton OpenAIProvider + MAX_AI_CHARS removal + ollama/lmstudio context_chars | 02bcbb9 | openai_provider.py, generic_openai_provider.py, ollama_provider.py, lmstudio_provider.py, classifier.py, requirements.txt |
| 3 | Registry-based get_provider(config: ProviderConfig) | 13eef37 | ai/__init__.py, provider_config.py |
| 4 | Promote 6 Wave-2 xfail stubs to passing | 209b156 | tests/test_ai_providers.py |
## What Was Built
### Task 1: ProviderConfig + PROVIDER_DEFAULTS + SUPPORTS_JSON_MODE
Created `backend/ai/provider_config.py` as a pure data module (no provider class imports):
- `ProviderConfig(BaseModel)` with `extra="forbid"`: provider_id (str, required), api_key (str, default ""), base_url (Optional[str], default None), model (str, default ""), context_chars (int, default 0 — sentinel meaning "use PROVIDER_DEFAULTS")
- `PROVIDER_DEFAULTS: dict[str, dict]` with 10 entries from RESEARCH.md Pattern 2: openai, anthropic, gemini, groq, xai, deepseek, openrouter, mistral, ollama, lmstudio — each with base_url, model, context_chars
- `SUPPORTS_JSON_MODE: dict[str, bool]` with gemini=False (OpenAI compat endpoint does not support `json_object` string form) and True for all others
### Task 2: GenericOpenAIProvider + Singleton OpenAIProvider + Removals
**backend/ai/openai_provider.py** refactored:
- `__init__` now accepts `context_chars: int` parameter
- `self._client = AsyncOpenAI(api_key=self._api_key, base_url=self._base_url)` stored as singleton in `__init__` (D-07)
- `def _client(self)` method deleted entirely
- `MAX_AI_CHARS = 8_000` constant deleted
- `def _truncate(self, text: str) -> str` added: returns text unchanged if len <= context_chars, otherwise `text[:head] + "\n[...truncated...]\n" + text[-tail:]` where head = int(context_chars*0.6), tail = context_chars - head (D-13)
- `classify()`, `suggest_topics()`, `health_check()` updated to use `self._client.chat.completions.create(...)` (no parentheses)
**backend/ai/generic_openai_provider.py** created:
- `class GenericOpenAIProvider(OpenAIProvider)` with `supports_json_mode` instance attribute
- `__init__` accepts `supports_json_mode: bool = True` kwarg, calls `super().__init__(...)`
- `classify()` and `suggest_topics()` conditionally add `response_format={"type":"json_object"}` when `supports_json_mode is True`; omit it for Gemini preset (D-01/D-02)
- Both methods import and call `parse_classification` / `parse_suggestions` from `ai.utils` — D-02 contract enforced via import line
- `health_check()` inherited from OpenAIProvider
**backend/ai/ollama_provider.py** and **lmstudio_provider.py**: Added `context_chars: int = 8000` parameter, passed through to `super().__init__()`.
**backend/services/classifier.py**: Removed `MAX_AI_CHARS = 8_000` constant and replaced `text[:MAX_AI_CHARS]` slices with `text` (truncation now inside provider via `_truncate()`).
**backend/requirements.txt**: `anthropic>=0.26` bumped to `anthropic>=0.95.0` (D-03 output_config support gate for Plan 03).
### Task 3: Registry-Based Factory
**backend/ai/__init__.py** rewritten:
- `_REGISTRY: dict[str, type[AIProvider]]` maps 10 provider_ids to classes
- `def get_provider(config: ProviderConfig) -> AIProvider` with typed signature (no raw dict)
- Resolves effective values from PROVIDER_DEFAULTS when config fields are empty/zero
- AnthropicProvider instantiated without base_url (current ctor contract; Plan 03 widens)
- GenericOpenAIProvider gets `supports_json_mode=SUPPORTS_JSON_MODE[config.provider_id]`
- Raises `ValueError(f"Unknown AI provider: {config.provider_id!r}")` for unknown ids
### Task 4: Six Wave-2 Tests Promoted
All 6 tests now pass without `@pytest.mark.xfail`:
1. **test_get_provider_typed**: Registry lookup returns GenericOpenAIProvider for groq/gemini; ValueError for unknown; supports_json_mode correct; _context_chars from PROVIDER_DEFAULTS
2. **test_client_singleton**: `AsyncOpenAI` class called exactly once per provider instance (mocked via patch)
3. **test_generic_openai_json_mode**: `response_format` present in kwargs when supports_json_mode=True; absent when False
4. **test_context_chars_truncation**: Provider with context_chars=100 truncates 500-char input with "[...truncated...]"
5. **test_smart_truncation**: 1000-char limit on 10000-char input → starts with 600 'H's, ends with 400 'T's
6. **test_gemini_fallback_to_parse_classification**: D-02 contract — `parse_classification` called with raw content AND `response_format` absent from API call kwargs
`test_anthropic_structured_output` remains xfail (Plan 07-03).
## Verification Results
- `grep -v '^#' backend/ai/openai_provider.py | grep -c 'MAX_AI_CHARS'` → 0
- `grep -v '^#' backend/services/classifier.py | grep -c 'MAX_AI_CHARS'` → 0
- `grep "from ai.utils import parse_classification" backend/ai/generic_openai_provider.py` → match found
- `grep -c "def parse_classification" backend/ai/utils.py` → 1 (file untouched)
- `backend/ai/__init__.py` contains `_REGISTRY` and `def get_provider(config: ProviderConfig)` and all 10 provider keys
- `requirements.txt` contains `anthropic>=0.95.0`
- Full test suite: **1 failed** (pre-existing test_extract_docx ModuleNotFoundError), **363 passed**, **15 xfailed**, **5 skipped** — no new failures; xfailed count down by 6
## Deviations from Plan
### Auto-fixed Issues
**1. [Rule 1 - Bug] ProviderConfig.context_chars default changed from 8000 to 0**
- **Found during:** Task 3 (test_get_provider_typed assertion failure)
- **Issue:** Task 1 spec said `context_chars: int = 8000` but Task 3 test asserts `result._context_chars == PROVIDER_DEFAULTS["groq"]["context_chars"]` (128000) when `ProviderConfig(provider_id="groq")` is created without specifying context_chars. With default=8000, `8000 or 128000 = 8000` (truthy short-circuit) — test fails.
- **Fix:** Changed `context_chars: int = 0` (sentinel meaning "unset — use PROVIDER_DEFAULTS in factory"). The factory's `config.context_chars or defaults["context_chars"]` then correctly resolves: `0 or 128000 = 128000`.
- **Files modified:** backend/ai/provider_config.py
- **Commit:** 13eef37
## Known Stubs
None — all plan goals implemented; no placeholders.
## Threat Flags
No new threat surface introduced. Changes are purely internal provider class refactoring and factory logic — no new API endpoints, no new DB access patterns, no new network paths. T-07-04 (empty api_key) mitigated: factory normalizes `api_key or "not-needed"` before passing to AsyncOpenAI constructor.
## Self-Check: PASSED
Files created/modified:
- [x] backend/ai/provider_config.py — FOUND (ProviderConfig + PROVIDER_DEFAULTS + SUPPORTS_JSON_MODE)
- [x] backend/ai/generic_openai_provider.py — FOUND (GenericOpenAIProvider class present)
- [x] backend/ai/openai_provider.py — FOUND (singleton _client, _truncate, no MAX_AI_CHARS)
- [x] backend/ai/ollama_provider.py — FOUND (context_chars param present)
- [x] backend/ai/lmstudio_provider.py — FOUND (context_chars param present)
- [x] backend/ai/__init__.py — FOUND (_REGISTRY and get_provider(config: ProviderConfig))
- [x] backend/services/classifier.py — FOUND (MAX_AI_CHARS removed)
- [x] backend/requirements.txt — FOUND (anthropic>=0.95.0)
- [x] backend/tests/test_ai_providers.py — FOUND (6 tests promoted, 1 xfail remaining)
Commits:
- [x] beb5b5e — feat(07-02): ProviderConfig Pydantic model + PROVIDER_DEFAULTS + SUPPORTS_JSON_MODE
- [x] 02bcbb9 — feat(07-02): singleton OpenAIProvider + GenericOpenAIProvider + MAX_AI_CHARS removal
- [x] 13eef37 — feat(07-02): registry-based get_provider(config: ProviderConfig) — D-06
- [x] 209b156 — test(07-02): promote 6 Wave-2 xfail tests to passing — D-01/D-02/D-07/D-12/D-13
+81 -27
View File
@@ -1,35 +1,89 @@
from ai.base import AIProvider, ClassificationResult """AI provider factory — registry-based O(1) lookup.
Usage:
from ai import get_provider
from ai.provider_config import ProviderConfig
config = ProviderConfig(provider_id="groq", api_key="sk-...")
provider = get_provider(config)
result = await provider.classify(text, topics, system_prompt)
"""
from __future__ import annotations
from ai.base import AIProvider
from ai.anthropic_provider import AnthropicProvider from ai.anthropic_provider import AnthropicProvider
from ai.openai_provider import OpenAIProvider from ai.openai_provider import OpenAIProvider
from ai.ollama_provider import OllamaProvider from ai.generic_openai_provider import GenericOpenAIProvider
from ai.lmstudio_provider import LMStudioProvider from ai.provider_config import ProviderConfig, PROVIDER_DEFAULTS, SUPPORTS_JSON_MODE
def get_provider(settings: dict) -> AIProvider: # Registry: maps provider_id → provider class
active = settings.get("active_provider", "lmstudio") # "openai" uses OpenAIProvider (no response_format override needed — plain OpenAI)
providers = settings.get("providers", {}) # "anthropic" uses AnthropicProvider (native output_config, no base_url ctor arg until Plan 03)
cfg = providers.get(active, {}) # All 8 OpenAI-compat vendors use GenericOpenAIProvider (D-16/D-17/D-18)
_REGISTRY: dict[str, type[AIProvider]] = {
"openai": OpenAIProvider,
"anthropic": AnthropicProvider,
"gemini": GenericOpenAIProvider,
"groq": GenericOpenAIProvider,
"xai": GenericOpenAIProvider,
"deepseek": GenericOpenAIProvider,
"openrouter": GenericOpenAIProvider,
"mistral": GenericOpenAIProvider,
"ollama": GenericOpenAIProvider,
"lmstudio": GenericOpenAIProvider,
}
if active == "anthropic":
return AnthropicProvider( def get_provider(config: ProviderConfig) -> AIProvider:
api_key=cfg.get("api_key", ""), """Instantiate and return an AI provider for the given ProviderConfig.
model=cfg.get("model", "claude-sonnet-4-6"),
Resolves defaults from PROVIDER_DEFAULTS when config fields are absent,
normalises an empty api_key to "not-needed" (OpenAI SDK 2.34+ rejects ""),
and sets supports_json_mode from the SUPPORTS_JSON_MODE lookup for
GenericOpenAIProvider instances.
Args:
config: A ProviderConfig with at minimum provider_id set.
Returns:
A fully-constructed AIProvider instance.
Raises:
ValueError: If config.provider_id is not in the registry.
"""
cls = _REGISTRY.get(config.provider_id)
if cls is None:
raise ValueError(f"Unknown AI provider: {config.provider_id!r}")
defaults = PROVIDER_DEFAULTS[config.provider_id]
# Resolve effective values — config fields take precedence over defaults
effective_api_key = config.api_key or "not-needed"
effective_model = config.model or defaults["model"]
effective_base_url = config.base_url if config.base_url is not None else defaults["base_url"]
effective_context_chars = config.context_chars or defaults["context_chars"]
if config.provider_id == "anthropic":
# AnthropicProvider does not accept base_url until Plan 03 refactors it;
# pass only the args its current __init__ accepts.
return cls(
api_key=effective_api_key,
model=effective_model,
) )
elif active == "openai": elif cls is GenericOpenAIProvider:
return OpenAIProvider( return cls(
api_key=cfg.get("api_key", ""), api_key=effective_api_key,
model=cfg.get("model", "gpt-4o"), model=effective_model,
base_url=cfg.get("base_url") or None, base_url=effective_base_url,
) context_chars=effective_context_chars,
elif active == "ollama": supports_json_mode=SUPPORTS_JSON_MODE[config.provider_id],
return OllamaProvider(
base_url=cfg.get("base_url", "http://host.docker.internal:11434"),
model=cfg.get("model", "llama3.2"),
)
elif active == "lmstudio":
return LMStudioProvider(
base_url=cfg.get("base_url", "http://host.docker.internal:1234"),
model=cfg.get("model", "gemma-4-e4b-it"),
) )
else: else:
raise ValueError(f"Unknown AI provider: {active}") # OpenAIProvider
return cls(
api_key=effective_api_key,
model=effective_model,
base_url=effective_base_url,
context_chars=effective_context_chars,
)
+103
View File
@@ -0,0 +1,103 @@
"""GenericOpenAIProvider — unified OpenAI-compatible provider for all 8 compat vendors.
Covers: Groq, xAI/Grok, DeepSeek, OpenRouter, Gemini-compat, Mistral-compat,
Ollama, LMStudio (D-16/D-17/D-18).
Key design decisions:
- Subclasses OpenAIProvider to inherit singleton _client and _truncate (D-07/D-13).
- Conditionally passes response_format={"type":"json_object"} based on
supports_json_mode flag (D-01) — Gemini preset sets this to False (D-02).
- Always parses the raw response with parse_classification / parse_suggestions
imported from ai.utils (D-02 last-resort fallback contract — NEVER redefine
locally; CLAUDE.md shared module map rule).
"""
from __future__ import annotations
from ai.openai_provider import OpenAIProvider
from ai.utils import parse_classification, parse_suggestions # D-02 contract
class GenericOpenAIProvider(OpenAIProvider):
"""OpenAI-compatible provider that enforces JSON mode on every call.
Named presets (Groq, xAI, DeepSeek, OpenRouter, Gemini-compat, Mistral-compat,
Ollama, LMStudio) are factory shortcuts in ai/__init__.py that pass the known
base_url default for each vendor.
supports_json_mode=False routes the provider through the parse_classification()
fallback path without sending response_format — used for Gemini (D-02).
"""
supports_json_mode: bool = True # class-level default; overridden per instance
def __init__(
self,
api_key: str,
model: str,
base_url: str | None,
context_chars: int = 8000,
supports_json_mode: bool = True,
):
super().__init__(
api_key=api_key,
model=model,
base_url=base_url,
context_chars=context_chars,
)
# Instance-level flag (may differ from class-level default)
self.supports_json_mode = supports_json_mode
async def classify(
self,
document_text: str,
existing_topics: list[str],
system_prompt: str,
):
topics_str = ", ".join(existing_topics) if existing_topics else "(none yet)"
user_msg = (
f"Existing topics: [{topics_str}]\n\n"
f"Document text:\n{self._truncate(document_text)}"
)
create_kwargs = dict(
model=self._model,
max_tokens=1024,
messages=[
{"role": "system", "content": system_prompt},
{"role": "user", "content": user_msg},
],
)
if self.supports_json_mode:
# D-01: enforce structured JSON output on all supporting providers
create_kwargs["response_format"] = {"type": "json_object"}
# else: Gemini preset — omit response_format, fall back to parse_classification()
response = await self._client.chat.completions.create(**create_kwargs)
raw = response.choices[0].message.content or ""
# D-02: parse_classification is the last-resort fallback — always called
return parse_classification(raw)
async def suggest_topics(
self,
document_text: str,
system_prompt: str,
) -> list[str]:
user_msg = (
"Suggest 3-5 topic names for this document. "
"Return ONLY valid JSON: {\"suggested_topics\": [\"topic1\", \"topic2\"]}\n\n"
f"Document text:\n{self._truncate(document_text)}"
)
create_kwargs = dict(
model=self._model,
max_tokens=256,
messages=[
{"role": "system", "content": system_prompt},
{"role": "user", "content": user_msg},
],
)
if self.supports_json_mode:
create_kwargs["response_format"] = {"type": "json_object"}
response = await self._client.chat.completions.create(**create_kwargs)
raw = response.choices[0].message.content or ""
# D-02: parse_suggestions is the last-resort fallback — always called
return parse_suggestions(raw)
# health_check is inherited from OpenAIProvider — no override needed
+7 -1
View File
@@ -2,9 +2,15 @@ from ai.openai_provider import OpenAIProvider
class LMStudioProvider(OpenAIProvider): class LMStudioProvider(OpenAIProvider):
def __init__(self, base_url: str = "http://host.docker.internal:1234", model: str = "gemma-4-e4b-it"): def __init__(
self,
base_url: str = "http://host.docker.internal:1234",
model: str = "gemma-4-e4b-it",
context_chars: int = 8000,
):
super().__init__( super().__init__(
api_key="lm-studio", api_key="lm-studio",
model=model, model=model,
base_url=base_url.rstrip("/") + "/v1", base_url=base_url.rstrip("/") + "/v1",
context_chars=context_chars,
) )
+7 -1
View File
@@ -2,9 +2,15 @@ from ai.openai_provider import OpenAIProvider
class OllamaProvider(OpenAIProvider): class OllamaProvider(OpenAIProvider):
def __init__(self, base_url: str = "http://host.docker.internal:11434", model: str = "llama3.2"): def __init__(
self,
base_url: str = "http://host.docker.internal:11434",
model: str = "llama3.2",
context_chars: int = 8000,
):
super().__init__( super().__init__(
api_key="ollama", api_key="ollama",
model=model, model=model,
base_url=base_url.rstrip("/") + "/v1", base_url=base_url.rstrip("/") + "/v1",
context_chars=context_chars,
) )
+32 -13
View File
@@ -1,18 +1,39 @@
from __future__ import annotations
from openai import AsyncOpenAI from openai import AsyncOpenAI
from ai.base import AIProvider, ClassificationResult from ai.base import AIProvider, ClassificationResult
from ai.utils import parse_classification, parse_suggestions from ai.utils import parse_classification, parse_suggestions
MAX_AI_CHARS = 8_000
class OpenAIProvider(AIProvider): class OpenAIProvider(AIProvider):
def __init__(self, api_key: str, model: str = "gpt-4o", base_url=None): # type: ignore[type-arg] def __init__(
self._api_key = api_key self,
api_key: str,
model: str = "gpt-4o",
base_url: str | None = None,
context_chars: int = 8000,
):
self._api_key = api_key or "not-needed"
self._model = model self._model = model
self._base_url = base_url self._base_url = base_url
self._context_chars = context_chars
# Singleton: created once in __init__, reused for all calls on this instance.
# Do NOT recreate per API call — AsyncOpenAI wraps an httpx.AsyncClient
# that maintains a connection pool; recreating per call destroys pool reuse
# and forces a new TLS handshake per request (D-07 / RESEARCH.md).
self._client = AsyncOpenAI(api_key=self._api_key, base_url=self._base_url)
def _client(self) -> AsyncOpenAI: def _truncate(self, text: str) -> str:
return AsyncOpenAI(api_key=self._api_key or "placeholder", base_url=self._base_url) """D-13 smart truncation: first 60% + last 40% of context window.
Captures both document introduction and conclusion, which carry the
most topic signal for long documents.
"""
if len(text) <= self._context_chars:
return text
head_len = int(self._context_chars * 0.6)
tail_len = self._context_chars - head_len
return text[:head_len] + "\n[...truncated...]\n" + text[-tail_len:]
async def classify( async def classify(
self, self,
@@ -23,9 +44,9 @@ class OpenAIProvider(AIProvider):
topics_str = ", ".join(existing_topics) if existing_topics else "(none yet)" topics_str = ", ".join(existing_topics) if existing_topics else "(none yet)"
user_msg = ( user_msg = (
f"Existing topics: [{topics_str}]\n\n" f"Existing topics: [{topics_str}]\n\n"
f"Document text:\n{document_text[:MAX_AI_CHARS]}" f"Document text:\n{self._truncate(document_text)}"
) )
response = await self._client().chat.completions.create( response = await self._client.chat.completions.create(
model=self._model, model=self._model,
max_tokens=1024, max_tokens=1024,
messages=[ messages=[
@@ -44,9 +65,9 @@ class OpenAIProvider(AIProvider):
user_msg = ( user_msg = (
"Suggest 3-5 topic names for this document. " "Suggest 3-5 topic names for this document. "
"Return ONLY valid JSON: {\"suggested_topics\": [\"topic1\", \"topic2\"]}\n\n" "Return ONLY valid JSON: {\"suggested_topics\": [\"topic1\", \"topic2\"]}\n\n"
f"Document text:\n{document_text[:MAX_AI_CHARS]}" f"Document text:\n{self._truncate(document_text)}"
) )
response = await self._client().chat.completions.create( response = await self._client.chat.completions.create(
model=self._model, model=self._model,
max_tokens=256, max_tokens=256,
messages=[ messages=[
@@ -59,7 +80,7 @@ class OpenAIProvider(AIProvider):
async def health_check(self) -> bool: async def health_check(self) -> bool:
try: try:
await self._client().chat.completions.create( await self._client.chat.completions.create(
model=self._model, model=self._model,
max_tokens=5, max_tokens=5,
messages=[{"role": "user", "content": "ping"}], messages=[{"role": "user", "content": "ping"}],
@@ -67,5 +88,3 @@ class OpenAIProvider(AIProvider):
return True return True
except Exception: except Exception:
return False return False
+107
View File
@@ -0,0 +1,107 @@
"""ProviderConfig Pydantic model and per-provider defaults.
Loaded by get_provider() in ai/__init__.py; populated by load_provider_config()
in services/ai_config.py.
This file is a pure data module — it does NOT import any provider class.
"""
from __future__ import annotations
from typing import Optional
from pydantic import BaseModel
class ProviderConfig(BaseModel):
"""Typed configuration for a single AI provider instance.
Fields:
provider_id: One of the keys in PROVIDER_DEFAULTS (e.g. "openai", "groq").
api_key: Decrypted API key; empty string for local providers (Ollama, LMStudio).
base_url: Override for the provider's base URL; None means use the default.
model: Model name; empty string means use the default from PROVIDER_DEFAULTS.
context_chars: Character budget for input truncation; 0 means use the default.
"""
model_config = {"extra": "forbid"}
provider_id: str
api_key: str = ""
base_url: Optional[str] = None
model: str = ""
context_chars: int = 0 # 0 means "unset — use PROVIDER_DEFAULTS in get_provider()"
# Named preset defaults for all 10 supported providers.
# Values are [ASSUMED] approximations based on well-known context window sizes.
# Admins can override all fields via the system_settings DB table (D-04).
PROVIDER_DEFAULTS: dict[str, dict] = {
"openai": {
"base_url": None,
"model": "gpt-4o",
"context_chars": 120_000,
},
"anthropic": {
"base_url": None,
"model": "claude-sonnet-4-6",
"context_chars": 180_000,
},
"gemini": {
"base_url": "https://generativelanguage.googleapis.com/v1beta/openai/",
"model": "gemini-2.0-flash",
"context_chars": 800_000,
},
"groq": {
"base_url": "https://api.groq.com/openai/v1",
"model": "llama-3.3-70b-versatile",
"context_chars": 128_000,
},
"xai": {
"base_url": "https://api.x.ai/v1",
"model": "grok-3-mini",
"context_chars": 128_000,
},
"deepseek": {
"base_url": "https://api.deepseek.com",
"model": "deepseek-chat",
"context_chars": 60_000,
},
"openrouter": {
"base_url": "https://openrouter.ai/api/v1",
"model": "anthropic/claude-3.5-sonnet",
"context_chars": 180_000,
},
"mistral": {
"base_url": "https://api.mistral.ai/v1",
"model": "mistral-large-latest",
"context_chars": 128_000,
},
"ollama": {
"base_url": "http://host.docker.internal:11434/v1",
"model": "llama3.2",
"context_chars": 8_000,
},
"lmstudio": {
"base_url": "http://host.docker.internal:1234/v1",
"model": "gemma-4-e4b-it",
"context_chars": 8_000,
},
}
# Whether the provider honours response_format={"type": "json_object"}.
# Gemini's OpenAI-compat endpoint does NOT support the string form (D-02/D-03).
# Ollama and LMStudio accept the parameter but some models ignore it — the
# GenericOpenAIProvider always wraps the raw response with parse_classification()
# regardless, so they are left as True (the parameter is still sent).
SUPPORTS_JSON_MODE: dict[str, bool] = {
"openai": True,
"anthropic": True,
"gemini": False,
"groq": True,
"xai": True,
"deepseek": True,
"openrouter": True,
"mistral": True,
"ollama": True,
"lmstudio": True,
}
+1 -1
View File
@@ -3,7 +3,7 @@ uvicorn[standard]>=0.29
python-multipart>=0.0.27 python-multipart>=0.0.27
pydantic-settings>=2.2 pydantic-settings>=2.2
pydantic[email]>=2.0 pydantic[email]>=2.0
anthropic>=0.26 anthropic>=0.95.0
openai>=1.30 openai>=1.30
PyMuPDF>=1.26.7 PyMuPDF>=1.26.7
python-docx>=1.1 python-docx>=1.1
+2 -4
View File
@@ -25,8 +25,6 @@ from db.models import Document
from services import storage from services import storage
from ai import get_provider from ai import get_provider
MAX_AI_CHARS = 8_000
_DEFAULT_SYSTEM_PROMPT = """You are a document classification assistant. When given a document's text content and a list of existing topics, you must: _DEFAULT_SYSTEM_PROMPT = """You are a document classification assistant. When given a document's text content and a list of existing topics, you must:
1. Assign the document to one or more relevant topics from the list. 1. Assign the document to one or more relevant topics from the list.
2. If no existing topics fit well, suggest new topic names. 2. If no existing topics fit well, suggest new topic names.
@@ -82,7 +80,7 @@ async def classify_document(
topic_names = [t["name"] for t in all_topics] topic_names = [t["name"] for t in all_topics]
text = meta.get("extracted_text", "") text = meta.get("extracted_text", "")
result = await provider.classify(text[:MAX_AI_CHARS], topic_names, system_prompt) result = await provider.classify(text, topic_names, system_prompt)
# Collect all topic names to persist (assigned + suggested) # Collect all topic names to persist (assigned + suggested)
all_new_names = set(result.suggested_new_topics) | set(result.topics) all_new_names = set(result.suggested_new_topics) | set(result.topics)
@@ -124,4 +122,4 @@ async def suggest_topics_for_document(
} }
provider = get_provider(_settings) provider = get_provider(_settings)
text = meta.get("extracted_text", "") text = meta.get("extracted_text", "")
return await provider.suggest_topics(text[:MAX_AI_CHARS], system_prompt) return await provider.suggest_topics(text, system_prompt)
+184 -33
View File
@@ -1,46 +1,197 @@
""" """
Wave 0 xfail stubs for Phase 7 AI provider tests. Tests for Phase 7 AI provider refactor.
Each function is a placeholder for a test that will be promoted to green Wave 2 (Plan 07-02) promotes: test_get_provider_typed, test_client_singleton,
in a later plan wave (per 07-VALIDATION.md per-task-verification map). test_generic_openai_json_mode, test_context_chars_truncation, test_smart_truncation,
test_gemini_fallback_to_parse_classification.
Stub policy (STATE.md decision: xfail(strict=False) for Wave 0): Remaining stubs (promoted in Plan 07-03): test_anthropic_structured_output.
- Body is a single pytest.xfail() call — no assertion code.
- strict=False so unexpected passes (xpass) never break CI.
""" """
import pytest import pytest
from unittest.mock import AsyncMock, MagicMock, patch
from ai.generic_openai_provider import GenericOpenAIProvider
from ai.openai_provider import OpenAIProvider
from ai.provider_config import ProviderConfig, PROVIDER_DEFAULTS
from ai.utils import parse_classification
@pytest.mark.xfail(strict=False, reason="Wave 0 stub — promoted in Plan 07-02") # ---------------------------------------------------------------------------
# Task 3: Registry-based get_provider(config: ProviderConfig) — D-06
# ---------------------------------------------------------------------------
def test_get_provider_typed():
"""get_provider() accepts a ProviderConfig and returns the correct provider class."""
from ai import get_provider
# Groq → GenericOpenAIProvider with supports_json_mode=True
config = ProviderConfig(provider_id="groq")
result = get_provider(config)
assert isinstance(result, GenericOpenAIProvider)
assert result.supports_json_mode is True
assert result._context_chars == PROVIDER_DEFAULTS["groq"]["context_chars"]
# Gemini → GenericOpenAIProvider with supports_json_mode=False (D-02)
config_gemini = ProviderConfig(provider_id="gemini")
result_gemini = get_provider(config_gemini)
assert isinstance(result_gemini, GenericOpenAIProvider)
assert result_gemini.supports_json_mode is False
# Unknown provider → ValueError
config_bogus = ProviderConfig(provider_id="bogus")
with pytest.raises(ValueError, match="Unknown AI provider"):
get_provider(config_bogus)
# ---------------------------------------------------------------------------
# Task 4: Singleton client lifecycle — D-07
# ---------------------------------------------------------------------------
@pytest.mark.asyncio
async def test_client_singleton():
"""AsyncOpenAI is instantiated exactly once per OpenAIProvider instance (D-07)."""
with patch("ai.openai_provider.AsyncOpenAI") as mock_cls:
# mock_cls() returns a mock instance — configure chat.completions.create
mock_instance = MagicMock()
mock_instance.chat = MagicMock()
mock_instance.chat.completions = MagicMock()
mock_instance.chat.completions.create = AsyncMock(
return_value=MagicMock(
choices=[MagicMock(message=MagicMock(content='{"assigned_topics":[],"new_topic_suggestions":[]}'))]
)
)
mock_cls.return_value = mock_instance
provider = OpenAIProvider(api_key="test-key", model="gpt-4o", base_url=None, context_chars=1000)
# Call classify twice on the same instance
await provider.classify("doc text", [], "sys")
await provider.classify("doc text", [], "sys")
# AsyncOpenAI class should have been called exactly once (in __init__)
assert mock_cls.call_count == 1
# ---------------------------------------------------------------------------
# Task 4: JSON-mode conditional — D-01
# ---------------------------------------------------------------------------
@pytest.mark.asyncio
async def test_generic_openai_json_mode(): async def test_generic_openai_json_mode():
pytest.xfail("not implemented yet — Plan 07-02") """GenericOpenAIProvider passes response_format only when supports_json_mode=True."""
synthetic_response = MagicMock(
choices=[MagicMock(message=MagicMock(
content='{"assigned_topics":["finance"],"new_topic_suggestions":[]}'
))]
)
# supports_json_mode=True → response_format present in call kwargs
with patch("ai.openai_provider.AsyncOpenAI") as mock_cls:
mock_instance = MagicMock()
mock_instance.chat.completions.create = AsyncMock(return_value=synthetic_response)
mock_cls.return_value = mock_instance
provider_json = GenericOpenAIProvider(
api_key="key", model="gpt-4o", base_url=None,
context_chars=1000, supports_json_mode=True
)
await provider_json.classify("text", [], "sys")
call_kwargs = mock_instance.chat.completions.create.call_args.kwargs
assert "response_format" in call_kwargs
assert call_kwargs["response_format"] == {"type": "json_object"}
# supports_json_mode=False → response_format absent from call kwargs (Gemini preset path)
with patch("ai.openai_provider.AsyncOpenAI") as mock_cls2:
mock_instance2 = MagicMock()
mock_instance2.chat.completions.create = AsyncMock(return_value=synthetic_response)
mock_cls2.return_value = mock_instance2
provider_no_json = GenericOpenAIProvider(
api_key="key", model="gemini-2.0-flash",
base_url="https://generativelanguage.googleapis.com/v1beta/openai/",
context_chars=1000, supports_json_mode=False
)
await provider_no_json.classify("text", [], "sys")
call_kwargs2 = mock_instance2.chat.completions.create.call_args.kwargs
assert "response_format" not in call_kwargs2
# ---------------------------------------------------------------------------
# Task 4: Smart truncation — D-12/D-13
# ---------------------------------------------------------------------------
def test_context_chars_truncation():
"""Provider with context_chars=100 truncates a 500-char input."""
provider = OpenAIProvider(api_key="", model="gpt-4o", base_url=None, context_chars=100)
long_text = "a" * 500
result = provider._truncate(long_text)
assert len(result) < 500
assert "[...truncated...]" in result
def test_smart_truncation():
"""_truncate uses 60% head + 40% tail of context_chars."""
provider = OpenAIProvider(api_key="", model="gpt-4o", base_url=None, context_chars=1000)
# Build a distinguishable input where head and tail chars differ
input_text = "H" * 5000 + "T" * 5000 # 10000 chars total
result = provider._truncate(input_text)
# head = int(1000 * 0.6) = 600, tail = 1000 - 600 = 400
assert result.startswith("H" * 600)
assert result.endswith("T" * 400)
assert "[...truncated...]" in result
# ---------------------------------------------------------------------------
# Task 4: Gemini fallback to parse_classification (D-02 contract enforcement)
# ---------------------------------------------------------------------------
@pytest.mark.asyncio
async def test_gemini_fallback_to_parse_classification():
"""D-02: GenericOpenAIProvider(supports_json_mode=False) calls parse_classification()
and does NOT send response_format to the API.
"""
raw_content = '{"assigned_topics":["x"],"new_topic_suggestions":[],"reasoning":"r"}'
with patch("ai.openai_provider.AsyncOpenAI") as mock_cls:
mock_create = AsyncMock(
return_value=MagicMock(
choices=[MagicMock(message=MagicMock(content=raw_content))]
)
)
mock_instance = MagicMock()
mock_instance.chat.completions.create = mock_create
mock_cls.return_value = mock_instance
# Wrap parse_classification so we can assert it was called
with patch(
"ai.generic_openai_provider.parse_classification",
wraps=parse_classification,
) as mock_parse:
provider = GenericOpenAIProvider(
api_key="",
model="gemini-2.0-flash",
base_url="https://generativelanguage.googleapis.com/v1beta/openai/",
context_chars=8000,
supports_json_mode=False,
)
result = await provider.classify("doc text", [], "sys")
# Result must be a valid ClassificationResult
assert result.topics == ["x"]
# parse_classification was called with the raw content (D-02 contract)
assert mock_parse.called
mock_parse.assert_called_once_with(raw_content)
# response_format must NOT have been sent to the API (Gemini preset path)
call_kwargs = mock_create.call_args.kwargs
assert "response_format" not in call_kwargs
# ---------------------------------------------------------------------------
# Stub: promoted in Plan 07-03
# ---------------------------------------------------------------------------
@pytest.mark.xfail(strict=False, reason="Wave 0 stub — promoted in Plan 07-03") @pytest.mark.xfail(strict=False, reason="Wave 0 stub — promoted in Plan 07-03")
async def test_anthropic_structured_output(): async def test_anthropic_structured_output():
pytest.xfail("not implemented yet — Plan 07-03") pytest.xfail("not implemented yet — Plan 07-03")
@pytest.mark.xfail(strict=False, reason="Wave 0 stub — promoted in Plan 07-02")
async def test_get_provider_typed():
pytest.xfail("not implemented yet — Plan 07-02")
@pytest.mark.xfail(strict=False, reason="Wave 0 stub — promoted in Plan 07-02")
async def test_client_singleton():
pytest.xfail("not implemented yet — Plan 07-02")
@pytest.mark.xfail(strict=False, reason="Wave 0 stub — promoted in Plan 07-03")
async def test_context_chars_truncation():
pytest.xfail("not implemented yet — Plan 07-03")
@pytest.mark.xfail(strict=False, reason="Wave 0 stub — promoted in Plan 07-03")
async def test_smart_truncation():
pytest.xfail("not implemented yet — Plan 07-03")
@pytest.mark.xfail(strict=False, reason="Wave 0 stub — promoted in Plan 07-02 Task 4 (D-02 Gemini fallback path)")
async def test_gemini_fallback_to_parse_classification():
pytest.xfail("not implemented yet — Plan 07-02")