chore: merge executor worktree (wave3-recovery/07-03)
This commit is contained in:
@@ -0,0 +1,168 @@
|
|||||||
|
---
|
||||||
|
phase: 07-redo-and-optimize-llm-integration
|
||||||
|
plan: "03"
|
||||||
|
subsystem: backend/ai-providers
|
||||||
|
tags:
|
||||||
|
- ai
|
||||||
|
- anthropic
|
||||||
|
- output_config
|
||||||
|
- structured-output
|
||||||
|
- singleton-client
|
||||||
|
- classifier-refactor
|
||||||
|
- db-driven-config
|
||||||
|
- wave-3
|
||||||
|
dependency_graph:
|
||||||
|
requires:
|
||||||
|
- "07-02 (ProviderConfig + GenericOpenAIProvider + registry — get_provider accepts ProviderConfig)"
|
||||||
|
- "07-01 (system_settings table + HKDF helpers + load_provider_config stub)"
|
||||||
|
provides:
|
||||||
|
- "AnthropicProvider singleton _client + output_config + _truncate (D-03/D-07/D-12/D-13)"
|
||||||
|
- "classifier.classify_document driven by load_provider_config with per-user/env fallback (D-04/D-06)"
|
||||||
|
- "ai_config.py stub removed — real ProviderConfig used everywhere"
|
||||||
|
- "3 previously-xfailed tests promoted to passing (test_anthropic_structured_output, test_api_key_encrypt_decrypt, test_anthropic_stop_reason_fallback)"
|
||||||
|
affects:
|
||||||
|
- "07-04 (Celery retry — classifier is now provider-agnostic; truncation inside providers)"
|
||||||
|
- "07-05 (Admin AI panel — ProviderConfig is the single config representation)"
|
||||||
|
tech_stack:
|
||||||
|
added: []
|
||||||
|
patterns:
|
||||||
|
- "output_config={\"format\": {\"type\": \"json_schema\", \"schema\": _SCHEMA}} for Anthropic constrained decoding (D-03)"
|
||||||
|
- "stop_reason guard: raw='' when stop_reason != 'end_turn' → graceful degradation (T-07-08)"
|
||||||
|
- "Uniform ctor signature: __init__(api_key, model, context_chars, base_url) across all providers"
|
||||||
|
- "load_provider_config(session) → ProviderConfig | None — DB authoritative, env fallback (D-04/D-15)"
|
||||||
|
- "Per-user override path: ProviderConfig from PROVIDER_DEFAULTS, empty api_key (T-07-06)"
|
||||||
|
key_files:
|
||||||
|
created: []
|
||||||
|
modified:
|
||||||
|
- backend/ai/anthropic_provider.py
|
||||||
|
- backend/ai/__init__.py
|
||||||
|
- backend/services/ai_config.py
|
||||||
|
- backend/services/classifier.py
|
||||||
|
- backend/tests/test_ai_providers.py
|
||||||
|
- backend/tests/test_ai_config.py
|
||||||
|
- backend/tests/test_classifier.py
|
||||||
|
decisions:
|
||||||
|
- "output_config={'format': {'type': 'json_schema', 'schema': ...}} chosen over tool_use — semantically correct for classification, no extra parsing layer needed"
|
||||||
|
- "base_url accepted in AnthropicProvider.__init__ for uniform factory signature but unused — SDK manages endpoint"
|
||||||
|
- "Per-user override path uses empty api_key — admin configures per-provider keys in system_settings; T-07-06 mitigated"
|
||||||
|
- "test_classifier.py tests updated to assert ProviderConfig properties instead of dict fields — D-06 contract enforced in tests"
|
||||||
|
- "test_anthropic_stop_reason_fallback added as extra regression guard for T-07-08"
|
||||||
|
metrics:
|
||||||
|
duration: "~25 minutes"
|
||||||
|
completed: "2026-06-04"
|
||||||
|
tasks_completed: 2
|
||||||
|
tasks_total: 2
|
||||||
|
files_created: 0
|
||||||
|
files_modified: 7
|
||||||
|
---
|
||||||
|
|
||||||
|
# Phase 7 Plan 03: Anthropic output_config + Classifier ProviderConfig Refactor Summary
|
||||||
|
|
||||||
|
AnthropicProvider refactored with singleton _client, output_config constrained-decoding structured output, and _truncate; classifier.py wired to load_provider_config(session) replacing inline dict construction; ai_config.py stub removed; 3 xfailed tests promoted.
|
||||||
|
|
||||||
|
## Tasks Completed
|
||||||
|
|
||||||
|
| Task | Description | Commit | Files |
|
||||||
|
|------|-------------|--------|-------|
|
||||||
|
| 1 | AnthropicProvider singleton + output_config + truncation + no MAX_AI_CHARS | efc177a | anthropic_provider.py, ai/__init__.py, tests/test_ai_providers.py |
|
||||||
|
| 2 | Classifier wired to load_provider_config + ai_config stub removed | 95c386f | services/ai_config.py, services/classifier.py, tests/test_ai_config.py, tests/test_classifier.py |
|
||||||
|
|
||||||
|
## What Was Built
|
||||||
|
|
||||||
|
### Task 1: AnthropicProvider Refactor
|
||||||
|
|
||||||
|
**backend/ai/anthropic_provider.py** rewritten:
|
||||||
|
- `MAX_AI_CHARS = 8_000` module constant deleted
|
||||||
|
- `def _client(self)` property method deleted
|
||||||
|
- `__init__` signature widened to `(api_key, model, context_chars, base_url=None)` — uniform factory contract
|
||||||
|
- `self._client = anthropic.AsyncAnthropic(api_key=self._api_key)` stored as singleton in `__init__` (D-07)
|
||||||
|
- `def _truncate(self, text)` added — identical 60/40 split pattern to OpenAIProvider (D-13)
|
||||||
|
- `_CLASSIFICATION_SCHEMA` and `_SUGGESTIONS_SCHEMA` module-level dicts added (required + additionalProperties=False)
|
||||||
|
- `classify()` passes `output_config={"format": {"type": "json_schema", "schema": _CLASSIFICATION_SCHEMA}}` (D-03)
|
||||||
|
- `suggest_topics()` passes `output_config={"format": {"type": "json_schema", "schema": _SUGGESTIONS_SCHEMA}}` (D-03)
|
||||||
|
- Graceful degradation: `stop_reason != "end_turn"` → `raw = ""` → `parse_classification("")` returns empty ClassificationResult (T-07-08)
|
||||||
|
- `health_check()` does NOT pass `output_config` — only verifies connectivity/auth
|
||||||
|
|
||||||
|
**backend/ai/__init__.py** updated:
|
||||||
|
- Anthropic branch now passes `context_chars=effective_context_chars, base_url=effective_base_url` to AnthropicProvider
|
||||||
|
- Comment updated to reflect widened ctor (Plan 03)
|
||||||
|
|
||||||
|
**backend/tests/test_ai_providers.py** updated:
|
||||||
|
- `test_anthropic_structured_output` promoted from xfail — patches `ai.anthropic_provider.anthropic.AsyncAnthropic`, asserts `output_config` kwarg present and matches `_CLASSIFICATION_SCHEMA`, asserts AsyncAnthropic constructed once (singleton)
|
||||||
|
- `test_anthropic_stop_reason_fallback` added — simulates `stop_reason="max_tokens"`, asserts empty ClassificationResult returned without crash (T-07-08 regression guard)
|
||||||
|
|
||||||
|
### Task 2: Classifier ProviderConfig Refactor + ai_config Stub Removal
|
||||||
|
|
||||||
|
**backend/services/ai_config.py** rewritten:
|
||||||
|
- `_ProviderConfigStub` class deleted entirely
|
||||||
|
- Module-level import: `from ai.provider_config import ProviderConfig, PROVIDER_DEFAULTS` (no more lazy import inside function body)
|
||||||
|
- `load_provider_config(session)` return type changed to `Optional[ProviderConfig]`
|
||||||
|
- `seed_system_settings_from_env` now uses `PROVIDER_DEFAULTS.get(provider_id, {}).get("context_chars", 8000)` for accurate default context_chars per provider
|
||||||
|
|
||||||
|
**backend/services/classifier.py** refactored:
|
||||||
|
- Imports: `from services.ai_config import load_provider_config` and `from ai.provider_config import ProviderConfig, PROVIDER_DEFAULTS`
|
||||||
|
- `classify_document`: inline `_settings = {...}` dict construction replaced with:
|
||||||
|
1. Per-user override path: `ProviderConfig(provider_id=ai_provider, ...)` from PROVIDER_DEFAULTS (api_key="" — T-07-06 mitigation documented)
|
||||||
|
2. System path: `await load_provider_config(session)` → None fallback to env-var defaults
|
||||||
|
3. `provider = get_provider(config)` — unchanged call signature
|
||||||
|
- `suggest_topics_for_document`: same load/override/fallback pattern applied
|
||||||
|
- No `text[:N]` slices remain — truncation fully delegated to provider `_truncate()`
|
||||||
|
|
||||||
|
**backend/tests/test_ai_config.py** promoted:
|
||||||
|
- `test_api_key_encrypt_decrypt`: round-trip smoke test (encrypt → decrypt → assert equality), domain salt isolation (different provider_id → different ciphertext), cross-domain decrypt raises `InvalidToken`
|
||||||
|
- `test_load_provider_config`: DB integration test skipped without `INTEGRATION=1` (psycopg guard)
|
||||||
|
|
||||||
|
**backend/tests/test_classifier.py** updated:
|
||||||
|
- `test_per_user_provider`: captures `ProviderConfig` instead of dict; asserts `config.provider_id == "openai"` and `config.model == "gpt-4o"` (D-06 contract enforced in tests)
|
||||||
|
- `test_default_provider_fallback`: patches `load_provider_config` to return None; captures ProviderConfig; asserts `config.provider_id == "ollama"` (env fallback path)
|
||||||
|
|
||||||
|
## Verification Results
|
||||||
|
|
||||||
|
- `grep -v '^#' backend/ai/anthropic_provider.py | grep -c 'MAX_AI_CHARS'` → 0
|
||||||
|
- `grep -c 'output_config=' backend/ai/anthropic_provider.py` → 3 (classify, suggest_topics, and schema comment)
|
||||||
|
- `grep -c '_ProviderConfigStub' backend/services/ai_config.py` → 0
|
||||||
|
- `grep -c 'await load_provider_config(' backend/services/classifier.py` → 2 (one per function)
|
||||||
|
- Full test suite: **1 failed** (pre-existing test_extract_docx ModuleNotFoundError), **366 passed**, **12 xfailed**, **6 skipped** — no new failures; xfailed count down by 3 from wave 2
|
||||||
|
|
||||||
|
## Deviations from Plan
|
||||||
|
|
||||||
|
### Auto-fixed Issues
|
||||||
|
|
||||||
|
**1. [Rule 1 - Bug] Classifier tests used dict-based interface**
|
||||||
|
- **Found during:** Task 2 (test_per_user_provider and test_default_provider_fallback asserted dict fields)
|
||||||
|
- **Issue:** After the refactor, `get_provider` receives a `ProviderConfig` object. The existing tests captured the dict argument and asserted `settings.get("active_provider")` etc. These would fail with `AttributeError` or wrong assertions on a ProviderConfig object.
|
||||||
|
- **Fix:** Updated `test_per_user_provider` to assert `config.provider_id == "openai"` and `config.model == "gpt-4o"`. Updated `test_default_provider_fallback` to patch `load_provider_config` returning None and assert `config.provider_id == "ollama"`.
|
||||||
|
- **Files modified:** backend/tests/test_classifier.py
|
||||||
|
- **Commit:** 95c386f
|
||||||
|
|
||||||
|
**2. [Rule 2 - Missing functionality] `seed_system_settings_from_env` used hardcoded context_chars=8000**
|
||||||
|
- **Found during:** Task 2 code review
|
||||||
|
- **Issue:** The original seed function always inserted `context_chars=8000` regardless of provider, which would insert the wrong default for providers like Anthropic (180,000) or Groq (128,000).
|
||||||
|
- **Fix:** Updated seed to use `PROVIDER_DEFAULTS.get(provider_id, {}).get("context_chars", 8000)` so the seeded row reflects the correct default for each provider.
|
||||||
|
- **Files modified:** backend/services/ai_config.py
|
||||||
|
- **Commit:** 95c386f
|
||||||
|
|
||||||
|
## Known Stubs
|
||||||
|
|
||||||
|
None — all plan goals implemented; no placeholders.
|
||||||
|
|
||||||
|
## Threat Flags
|
||||||
|
|
||||||
|
No new threat surface introduced. Changes are internal service-layer refactoring:
|
||||||
|
- T-07-06 (per-user api_key isolation) mitigated: classifier per-user override path uses `api_key=""` — documented in comment
|
||||||
|
- T-07-08 (Anthropic stop_reason degradation) mitigated: classify() falls back to `parse_classification("")` when `stop_reason != "end_turn"`
|
||||||
|
|
||||||
|
## Self-Check: PASSED
|
||||||
|
|
||||||
|
Files modified:
|
||||||
|
- [x] backend/ai/anthropic_provider.py — FOUND (singleton _client, output_config, _truncate, no MAX_AI_CHARS)
|
||||||
|
- [x] backend/ai/__init__.py — FOUND (context_chars + base_url passed to AnthropicProvider)
|
||||||
|
- [x] backend/services/ai_config.py — FOUND (ProviderConfig import at module level, no _ProviderConfigStub)
|
||||||
|
- [x] backend/services/classifier.py — FOUND (load_provider_config + ProviderConfig construction)
|
||||||
|
- [x] backend/tests/test_ai_providers.py — FOUND (test_anthropic_structured_output promoted)
|
||||||
|
- [x] backend/tests/test_ai_config.py — FOUND (test_api_key_encrypt_decrypt promoted)
|
||||||
|
- [x] backend/tests/test_classifier.py — FOUND (per_user_provider + default_fallback tests updated)
|
||||||
|
|
||||||
|
Commits:
|
||||||
|
- [x] efc177a — feat(07-03): AnthropicProvider singleton + output_config + truncation — D-03/D-07/D-12/D-13
|
||||||
|
- [x] 95c386f — feat(07-03): classifier wired to load_provider_config + ai_config stub removed — D-04/D-06
|
||||||
@@ -19,7 +19,7 @@ from ai.provider_config import ProviderConfig, PROVIDER_DEFAULTS, SUPPORTS_JSON_
|
|||||||
|
|
||||||
# Registry: maps provider_id → provider class
|
# Registry: maps provider_id → provider class
|
||||||
# "openai" uses OpenAIProvider (no response_format override needed — plain OpenAI)
|
# "openai" uses OpenAIProvider (no response_format override needed — plain OpenAI)
|
||||||
# "anthropic" uses AnthropicProvider (native output_config, no base_url ctor arg until Plan 03)
|
# "anthropic" uses AnthropicProvider (native output_config, Plan 03 widened ctor accepts context_chars+base_url)
|
||||||
# All 8 OpenAI-compat vendors use GenericOpenAIProvider (D-16/D-17/D-18)
|
# All 8 OpenAI-compat vendors use GenericOpenAIProvider (D-16/D-17/D-18)
|
||||||
_REGISTRY: dict[str, type[AIProvider]] = {
|
_REGISTRY: dict[str, type[AIProvider]] = {
|
||||||
"openai": OpenAIProvider,
|
"openai": OpenAIProvider,
|
||||||
@@ -65,11 +65,13 @@ def get_provider(config: ProviderConfig) -> AIProvider:
|
|||||||
effective_context_chars = config.context_chars or defaults["context_chars"]
|
effective_context_chars = config.context_chars or defaults["context_chars"]
|
||||||
|
|
||||||
if config.provider_id == "anthropic":
|
if config.provider_id == "anthropic":
|
||||||
# AnthropicProvider does not accept base_url until Plan 03 refactors it;
|
# AnthropicProvider accepts context_chars and base_url (Plan 03 widened ctor).
|
||||||
# pass only the args its current __init__ accepts.
|
# base_url is accepted for uniform factory signature but unused by the SDK.
|
||||||
return cls(
|
return cls(
|
||||||
api_key=effective_api_key,
|
api_key=effective_api_key,
|
||||||
model=effective_model,
|
model=effective_model,
|
||||||
|
context_chars=effective_context_chars,
|
||||||
|
base_url=effective_base_url,
|
||||||
)
|
)
|
||||||
elif cls is GenericOpenAIProvider:
|
elif cls is GenericOpenAIProvider:
|
||||||
return cls(
|
return cls(
|
||||||
|
|||||||
@@ -1,17 +1,83 @@
|
|||||||
|
"""Anthropic AI provider — singleton client, output_config structured output, smart truncation.
|
||||||
|
|
||||||
|
D-03: Uses output_config={"format": {"type": "json_schema", "schema": ...}} with constrained
|
||||||
|
decoding available in anthropic SDK >=0.95.0 (GA, no beta headers needed).
|
||||||
|
D-07: self._client = AsyncAnthropic(...) created once in __init__ and reused — never recreated
|
||||||
|
per API call to preserve the httpx connection pool.
|
||||||
|
D-12/D-13: Global char constant removed; uses self._context_chars with 60/40 smart truncation.
|
||||||
|
|
||||||
|
Security: api_key is accepted from the caller (loaded from system_settings by ai_config.py
|
||||||
|
and decrypted before being passed here). The key is never stored beyond this instance's
|
||||||
|
lifetime. T-07-06 mitigated: this class never reads the api_key from env vars directly.
|
||||||
|
"""
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
import anthropic
|
import anthropic
|
||||||
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
|
# ── Output schemas for constrained decoding (D-03 / RESEARCH.md) ────────────────────────────
|
||||||
|
# additionalProperties=False required by Anthropic output_config grammar.
|
||||||
|
# "reasoning" is intentionally absent from "required" so legacy prompts that don't include it
|
||||||
|
# still produce valid JSON (Anthropic will emit it because it is declared in properties, but
|
||||||
|
# we do not enforce it in the schema to avoid refusal on minimal responses).
|
||||||
|
|
||||||
|
_CLASSIFICATION_SCHEMA: dict = {
|
||||||
|
"type": "object",
|
||||||
|
"properties": {
|
||||||
|
"assigned_topics": {"type": "array", "items": {"type": "string"}},
|
||||||
|
"new_topic_suggestions": {"type": "array", "items": {"type": "string"}},
|
||||||
|
"reasoning": {"type": "string"},
|
||||||
|
},
|
||||||
|
"required": ["assigned_topics", "new_topic_suggestions"],
|
||||||
|
"additionalProperties": False,
|
||||||
|
}
|
||||||
|
|
||||||
|
_SUGGESTIONS_SCHEMA: dict = {
|
||||||
|
"type": "object",
|
||||||
|
"properties": {
|
||||||
|
"suggested_topics": {"type": "array", "items": {"type": "string"}},
|
||||||
|
},
|
||||||
|
"required": ["suggested_topics"],
|
||||||
|
"additionalProperties": False,
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
class AnthropicProvider(AIProvider):
|
class AnthropicProvider(AIProvider):
|
||||||
def __init__(self, api_key: str, model: str = "claude-sonnet-4-6"):
|
"""Anthropic Claude provider with singleton client and output_config structured output.
|
||||||
|
|
||||||
|
Constructor signature matches the uniform factory contract in ai/__init__.py:
|
||||||
|
api_key, model, context_chars, base_url (accepted but unused — Anthropic SDK
|
||||||
|
manages the endpoint; widened so get_provider() can call all providers uniformly).
|
||||||
|
"""
|
||||||
|
|
||||||
|
def __init__(
|
||||||
|
self,
|
||||||
|
api_key: str,
|
||||||
|
model: str = "claude-sonnet-4-6",
|
||||||
|
context_chars: int = 180_000,
|
||||||
|
base_url: str | None = None, # accepted for uniform factory signature; unused
|
||||||
|
):
|
||||||
self._api_key = api_key
|
self._api_key = api_key
|
||||||
self._model = model
|
self._model = model
|
||||||
|
self._context_chars = context_chars
|
||||||
|
# Singleton: created once in __init__, reused for all calls on this instance.
|
||||||
|
# Do NOT recreate per API call — AsyncAnthropic 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 = anthropic.AsyncAnthropic(api_key=self._api_key)
|
||||||
|
|
||||||
def _client(self):
|
def _truncate(self, text: str) -> str:
|
||||||
return anthropic.AsyncAnthropic(api_key=self._api_key)
|
"""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,
|
||||||
@@ -22,16 +88,23 @@ class AnthropicProvider(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)}"
|
||||||
)
|
)
|
||||||
client = self._client()
|
response = await self._client.messages.create(
|
||||||
response = await client.messages.create(
|
|
||||||
model=self._model,
|
model=self._model,
|
||||||
max_tokens=1024,
|
max_tokens=1024,
|
||||||
system=system_prompt,
|
system=system_prompt,
|
||||||
messages=[{"role": "user", "content": user_msg}],
|
messages=[{"role": "user", "content": user_msg}],
|
||||||
|
output_config={"format": {"type": "json_schema", "schema": _CLASSIFICATION_SCHEMA}},
|
||||||
)
|
)
|
||||||
|
# Graceful degradation (T-07-08): when stop_reason is "refusal" or "max_tokens"
|
||||||
|
# the constrained decoding did not complete — fall back to parse_classification("")
|
||||||
|
# which returns an empty ClassificationResult rather than raising an exception.
|
||||||
|
stop_reason = getattr(response, "stop_reason", "end_turn")
|
||||||
|
if response.content and stop_reason == "end_turn":
|
||||||
raw = response.content[0].text
|
raw = response.content[0].text
|
||||||
|
else:
|
||||||
|
raw = ""
|
||||||
return parse_classification(raw)
|
return parse_classification(raw)
|
||||||
|
|
||||||
async def suggest_topics(
|
async def suggest_topics(
|
||||||
@@ -42,28 +115,34 @@ class AnthropicProvider(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)}"
|
||||||
)
|
)
|
||||||
client = self._client()
|
response = await self._client.messages.create(
|
||||||
response = await client.messages.create(
|
|
||||||
model=self._model,
|
model=self._model,
|
||||||
max_tokens=256,
|
max_tokens=256,
|
||||||
system=system_prompt,
|
system=system_prompt,
|
||||||
messages=[{"role": "user", "content": user_msg}],
|
messages=[{"role": "user", "content": user_msg}],
|
||||||
|
output_config={"format": {"type": "json_schema", "schema": _SUGGESTIONS_SCHEMA}},
|
||||||
)
|
)
|
||||||
|
stop_reason = getattr(response, "stop_reason", "end_turn")
|
||||||
|
if response.content and stop_reason == "end_turn":
|
||||||
raw = response.content[0].text
|
raw = response.content[0].text
|
||||||
|
else:
|
||||||
|
raw = ""
|
||||||
return parse_suggestions(raw)
|
return parse_suggestions(raw)
|
||||||
|
|
||||||
async def health_check(self) -> bool:
|
async def health_check(self) -> bool:
|
||||||
|
"""Verify API key validity and connectivity by sending a minimal message.
|
||||||
|
|
||||||
|
Does NOT pass output_config — the response shape does not matter here;
|
||||||
|
this only confirms the api_key and network path are working.
|
||||||
|
"""
|
||||||
try:
|
try:
|
||||||
client = self._client()
|
await self._client.messages.create(
|
||||||
await client.messages.create(
|
|
||||||
model=self._model,
|
model=self._model,
|
||||||
max_tokens=5,
|
max_tokens=8,
|
||||||
messages=[{"role": "user", "content": "ping"}],
|
messages=[{"role": "user", "content": "ping"}],
|
||||||
)
|
)
|
||||||
return True
|
return True
|
||||||
except Exception:
|
except Exception:
|
||||||
return False
|
return False
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
@@ -22,38 +22,21 @@ from __future__ import annotations
|
|||||||
|
|
||||||
import base64
|
import base64
|
||||||
import logging
|
import logging
|
||||||
from typing import Optional, TYPE_CHECKING
|
from typing import Optional
|
||||||
|
|
||||||
import structlog
|
import structlog
|
||||||
from cryptography.fernet import Fernet
|
from cryptography.fernet import Fernet
|
||||||
from cryptography.hazmat.primitives import hashes
|
from cryptography.hazmat.primitives import hashes
|
||||||
from cryptography.hazmat.primitives.kdf.hkdf import HKDF
|
from cryptography.hazmat.primitives.kdf.hkdf import HKDF
|
||||||
from pydantic import BaseModel
|
|
||||||
from sqlalchemy import select
|
from sqlalchemy import select
|
||||||
from sqlalchemy.ext.asyncio import AsyncSession
|
from sqlalchemy.ext.asyncio import AsyncSession
|
||||||
|
|
||||||
|
from ai.provider_config import ProviderConfig, PROVIDER_DEFAULTS
|
||||||
from config import settings
|
from config import settings
|
||||||
|
|
||||||
# ProviderConfig redefined in ai/provider_config.py during Plan 02 —
|
|
||||||
# load_provider_config will re-import and return that class once Plan 02 lands.
|
|
||||||
# The stub below is removed in Plan 03 once the classifier consumes the real class.
|
|
||||||
|
|
||||||
|
|
||||||
logger = structlog.get_logger(__name__)
|
logger = structlog.get_logger(__name__)
|
||||||
|
|
||||||
|
|
||||||
# ── Stub ProviderConfig for Plan 01 (replaced by ai/provider_config.py in Plan 02) ──
|
|
||||||
|
|
||||||
class _ProviderConfigStub(BaseModel):
|
|
||||||
"""Minimal provider config placeholder until Plan 02 creates ai/provider_config.py."""
|
|
||||||
|
|
||||||
provider_id: str
|
|
||||||
api_key: str = ""
|
|
||||||
base_url: Optional[str] = None
|
|
||||||
model: str = ""
|
|
||||||
context_chars: int = 8000
|
|
||||||
|
|
||||||
|
|
||||||
# ── HKDF key derivation ───────────────────────────────────────────────────────
|
# ── HKDF key derivation ───────────────────────────────────────────────────────
|
||||||
|
|
||||||
def _derive_ai_settings_key(master_key: bytes, provider_id: str) -> Fernet:
|
def _derive_ai_settings_key(master_key: bytes, provider_id: str) -> Fernet:
|
||||||
@@ -125,31 +108,19 @@ def decrypt_api_key(master_key: bytes, provider_id: str, api_key_enc: str) -> st
|
|||||||
|
|
||||||
# ── Provider config loader ────────────────────────────────────────────────────
|
# ── Provider config loader ────────────────────────────────────────────────────
|
||||||
|
|
||||||
async def load_provider_config(session: AsyncSession) -> Optional[_ProviderConfigStub]:
|
async def load_provider_config(session: AsyncSession) -> Optional[ProviderConfig]:
|
||||||
"""Load the active AI provider config from the system_settings table.
|
"""Load the active AI provider config from the system_settings table.
|
||||||
|
|
||||||
Returns a _ProviderConfigStub built from the row where is_active=True,
|
Returns a ProviderConfig built from the row where is_active=True,
|
||||||
decrypting api_key_enc when present. Returns None when no active row exists.
|
decrypting api_key_enc when present. Returns None when no active row exists.
|
||||||
|
|
||||||
Note: In Plan 02, this function will try to import and return ProviderConfig
|
|
||||||
from ai/provider_config.py instead of the stub. The import is done lazily
|
|
||||||
inside the function body so that Plan 01 does not depend on files that don't
|
|
||||||
exist yet.
|
|
||||||
|
|
||||||
Args:
|
Args:
|
||||||
session: An open AsyncSession.
|
session: An open AsyncSession.
|
||||||
|
|
||||||
Returns:
|
Returns:
|
||||||
A _ProviderConfigStub (or the real ProviderConfig from Plan 02+) if an
|
A ProviderConfig if an active row exists; None if the table is empty or
|
||||||
active row exists; None if the table is empty or no row is marked active.
|
no row is marked active.
|
||||||
"""
|
"""
|
||||||
# Lazy import: try to use the real ProviderConfig once Plan 02 lands
|
|
||||||
try:
|
|
||||||
from ai.provider_config import ProviderConfig as _RealProviderConfig # type: ignore[import]
|
|
||||||
config_cls = _RealProviderConfig
|
|
||||||
except ImportError:
|
|
||||||
config_cls = _ProviderConfigStub # type: ignore[assignment]
|
|
||||||
|
|
||||||
from db.models import SystemSettings # local import to avoid circular deps
|
from db.models import SystemSettings # local import to avoid circular deps
|
||||||
|
|
||||||
stmt = select(SystemSettings).where(SystemSettings.is_active.is_(True))
|
stmt = select(SystemSettings).where(SystemSettings.is_active.is_(True))
|
||||||
@@ -171,7 +142,7 @@ async def load_provider_config(session: AsyncSession) -> Optional[_ProviderConfi
|
|||||||
provider_id=row.provider_id,
|
provider_id=row.provider_id,
|
||||||
)
|
)
|
||||||
|
|
||||||
return config_cls(
|
return ProviderConfig(
|
||||||
provider_id=row.provider_id,
|
provider_id=row.provider_id,
|
||||||
api_key=api_key,
|
api_key=api_key,
|
||||||
base_url=row.base_url,
|
base_url=row.base_url,
|
||||||
@@ -208,11 +179,14 @@ async def seed_system_settings_from_env(session: AsyncSession) -> None:
|
|||||||
# Row already exists — never overwrite (idempotent)
|
# Row already exists — never overwrite (idempotent)
|
||||||
return
|
return
|
||||||
|
|
||||||
|
# Use PROVIDER_DEFAULTS context_chars for the provider, fallback to 8000
|
||||||
|
context_chars = PROVIDER_DEFAULTS.get(provider_id, {}).get("context_chars", 8000)
|
||||||
|
|
||||||
# Insert default row with no API key (local providers like Ollama don't need one)
|
# Insert default row with no API key (local providers like Ollama don't need one)
|
||||||
row = SystemSettings(
|
row = SystemSettings(
|
||||||
provider_id=provider_id,
|
provider_id=provider_id,
|
||||||
model_name=model_name,
|
model_name=model_name,
|
||||||
context_chars=8000,
|
context_chars=context_chars,
|
||||||
is_active=True,
|
is_active=True,
|
||||||
api_key_enc=None,
|
api_key_enc=None,
|
||||||
base_url=None,
|
base_url=None,
|
||||||
|
|||||||
@@ -13,6 +13,13 @@ the user's namespace via create_topic(user_id=doc.user_id) (D-11).
|
|||||||
Updated in Plan 03-04: classify_document and suggest_topics_for_document now accept
|
Updated in Plan 03-04: classify_document and suggest_topics_for_document now accept
|
||||||
ai_provider and ai_model kwargs. No longer calls storage.load_settings(). Provider
|
ai_provider and ai_model kwargs. No longer calls storage.load_settings(). Provider
|
||||||
resolved via get_provider() using per-user settings from DB (D-14, D-15).
|
resolved via get_provider() using per-user settings from DB (D-14, D-15).
|
||||||
|
|
||||||
|
Updated in Plan 07-03: Provider resolved via load_provider_config(session) reading
|
||||||
|
from the system_settings DB table (D-04/D-06). Per-user override still honoured:
|
||||||
|
when ai_provider is non-None, a ProviderConfig is constructed from PROVIDER_DEFAULTS
|
||||||
|
for that provider (note: per-user override does not carry an api_key — admin must
|
||||||
|
configure each provider's key in system_settings). Truncation delegated to provider
|
||||||
|
_truncate() — no more inline text slicing in the classifier (D-12/D-13).
|
||||||
"""
|
"""
|
||||||
from __future__ import annotations
|
from __future__ import annotations
|
||||||
|
|
||||||
@@ -23,7 +30,9 @@ from sqlalchemy.ext.asyncio import AsyncSession
|
|||||||
from config import settings as app_settings
|
from config import settings as app_settings
|
||||||
from db.models import Document
|
from db.models import Document
|
||||||
from services import storage
|
from services import storage
|
||||||
|
from services.ai_config import load_provider_config
|
||||||
from ai import get_provider
|
from ai import get_provider
|
||||||
|
from ai.provider_config import ProviderConfig, PROVIDER_DEFAULTS
|
||||||
|
|
||||||
_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.
|
||||||
@@ -47,19 +56,45 @@ async def classify_document(
|
|||||||
|
|
||||||
ai_provider and ai_model come from the document owner's User record (D-14).
|
ai_provider and ai_model come from the document owner's User record (D-14).
|
||||||
Falls back to app_settings.default_ai_provider / default_ai_model when None (D-15).
|
Falls back to app_settings.default_ai_provider / default_ai_model when None (D-15).
|
||||||
|
Provider config loaded from system_settings DB table via load_provider_config() (D-04).
|
||||||
"""
|
"""
|
||||||
meta = await storage.get_metadata(session, doc_id)
|
meta = await storage.get_metadata(session, doc_id)
|
||||||
if meta is None:
|
if meta is None:
|
||||||
raise ValueError(f"Document {doc_id} not found")
|
raise ValueError(f"Document {doc_id} not found")
|
||||||
|
|
||||||
_ai_provider = ai_provider or app_settings.default_ai_provider
|
|
||||||
_ai_model = ai_model or app_settings.default_ai_model
|
|
||||||
system_prompt = app_settings.system_prompt or _DEFAULT_SYSTEM_PROMPT
|
system_prompt = app_settings.system_prompt or _DEFAULT_SYSTEM_PROMPT
|
||||||
_settings = {
|
|
||||||
"active_provider": _ai_provider,
|
# ── Provider resolution (D-04/D-06) ─────────────────────────────────────
|
||||||
"providers": {_ai_provider: {"model": _ai_model}},
|
if ai_provider is not None:
|
||||||
}
|
# Per-user override path: build a ProviderConfig from PROVIDER_DEFAULTS.
|
||||||
provider = get_provider(_settings)
|
# per-user override does not carry an api_key — admin must configure each
|
||||||
|
# provider's key in system_settings. When the per-user override selects a
|
||||||
|
# different provider than the active system provider, api_key stays empty
|
||||||
|
# and get_provider() normalises it to "not-needed".
|
||||||
|
config = ProviderConfig(
|
||||||
|
provider_id=ai_provider,
|
||||||
|
model=ai_model or PROVIDER_DEFAULTS.get(ai_provider, {}).get("model", ""),
|
||||||
|
api_key="",
|
||||||
|
base_url=None,
|
||||||
|
context_chars=PROVIDER_DEFAULTS.get(ai_provider, {}).get("context_chars", 8000),
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
# System provider path: load from DB (D-04)
|
||||||
|
config = await load_provider_config(session)
|
||||||
|
if config is None:
|
||||||
|
# No active row in system_settings — fall back to env-var defaults (D-15)
|
||||||
|
fallback_provider = app_settings.default_ai_provider
|
||||||
|
config = ProviderConfig(
|
||||||
|
provider_id=fallback_provider,
|
||||||
|
model=app_settings.default_ai_model,
|
||||||
|
api_key="",
|
||||||
|
base_url=None,
|
||||||
|
context_chars=PROVIDER_DEFAULTS.get(fallback_provider, {}).get(
|
||||||
|
"context_chars", 8000
|
||||||
|
),
|
||||||
|
)
|
||||||
|
|
||||||
|
provider = get_provider(config)
|
||||||
|
|
||||||
# Load the Document ORM object to get the owner's user_id (D-11, D-17)
|
# Load the Document ORM object to get the owner's user_id (D-11, D-17)
|
||||||
try:
|
try:
|
||||||
@@ -80,6 +115,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", "")
|
||||||
|
# Truncation is performed inside provider.classify() via provider._truncate(). (D-12/D-13)
|
||||||
result = await provider.classify(text, 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)
|
||||||
@@ -108,18 +144,38 @@ async def suggest_topics_for_document(
|
|||||||
|
|
||||||
ai_provider and ai_model come from the document owner's User record (D-14).
|
ai_provider and ai_model come from the document owner's User record (D-14).
|
||||||
Falls back to app_settings.default_ai_provider / default_ai_model when None (D-15).
|
Falls back to app_settings.default_ai_provider / default_ai_model when None (D-15).
|
||||||
|
Provider config loaded from system_settings DB table via load_provider_config() (D-04).
|
||||||
"""
|
"""
|
||||||
meta = await storage.get_metadata(session, doc_id)
|
meta = await storage.get_metadata(session, doc_id)
|
||||||
if meta is None:
|
if meta is None:
|
||||||
raise ValueError(f"Document {doc_id} not found")
|
raise ValueError(f"Document {doc_id} not found")
|
||||||
|
|
||||||
_ai_provider = ai_provider or app_settings.default_ai_provider
|
|
||||||
_ai_model = ai_model or app_settings.default_ai_model
|
|
||||||
system_prompt = app_settings.system_prompt or _DEFAULT_SYSTEM_PROMPT
|
system_prompt = app_settings.system_prompt or _DEFAULT_SYSTEM_PROMPT
|
||||||
_settings = {
|
|
||||||
"active_provider": _ai_provider,
|
# ── Provider resolution (D-04/D-06) ─────────────────────────────────────
|
||||||
"providers": {_ai_provider: {"model": _ai_model}},
|
if ai_provider is not None:
|
||||||
}
|
config = ProviderConfig(
|
||||||
provider = get_provider(_settings)
|
provider_id=ai_provider,
|
||||||
|
model=ai_model or PROVIDER_DEFAULTS.get(ai_provider, {}).get("model", ""),
|
||||||
|
api_key="",
|
||||||
|
base_url=None,
|
||||||
|
context_chars=PROVIDER_DEFAULTS.get(ai_provider, {}).get("context_chars", 8000),
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
config = await load_provider_config(session)
|
||||||
|
if config is None:
|
||||||
|
fallback_provider = app_settings.default_ai_provider
|
||||||
|
config = ProviderConfig(
|
||||||
|
provider_id=fallback_provider,
|
||||||
|
model=app_settings.default_ai_model,
|
||||||
|
api_key="",
|
||||||
|
base_url=None,
|
||||||
|
context_chars=PROVIDER_DEFAULTS.get(fallback_provider, {}).get(
|
||||||
|
"context_chars", 8000
|
||||||
|
),
|
||||||
|
)
|
||||||
|
|
||||||
|
provider = get_provider(config)
|
||||||
text = meta.get("extracted_text", "")
|
text = meta.get("extracted_text", "")
|
||||||
|
# Truncation is performed inside provider.suggest_topics() via provider._truncate(). (D-12/D-13)
|
||||||
return await provider.suggest_topics(text, system_prompt)
|
return await provider.suggest_topics(text, system_prompt)
|
||||||
|
|||||||
@@ -1,25 +1,100 @@
|
|||||||
"""
|
"""
|
||||||
Wave 0 xfail stubs for Phase 7 AI config service tests.
|
Tests for Phase 7 AI config service.
|
||||||
|
|
||||||
Covers:
|
Covers:
|
||||||
- D-04: load_provider_config() reads from system_settings table
|
- D-04: load_provider_config() reads from system_settings table
|
||||||
- D-05: API key HKDF encryption round-trip
|
- D-05: API key HKDF encryption round-trip
|
||||||
|
|
||||||
Each function is a placeholder to be promoted in Plan 07-01 (green gate) once
|
Plan 07-03 promotes: test_api_key_encrypt_decrypt, test_load_provider_config.
|
||||||
the services/ai_config.py module and system_settings table exist.
|
|
||||||
|
|
||||||
Stub policy (STATE.md decision: xfail(strict=False) for Wave 0):
|
|
||||||
- Body is a single pytest.xfail() call — no assertion code.
|
|
||||||
- strict=False so unexpected passes (xpass) never break CI.
|
|
||||||
"""
|
"""
|
||||||
|
import os
|
||||||
import pytest
|
import pytest
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.xfail(strict=False, reason="Wave 0 stub — promoted in Plan 07-01")
|
# ---------------------------------------------------------------------------
|
||||||
async def test_load_provider_config():
|
# test_api_key_encrypt_decrypt — D-05 HKDF encryption round-trip
|
||||||
pytest.xfail("not implemented yet — Plan 07-01")
|
# ---------------------------------------------------------------------------
|
||||||
|
|
||||||
|
def test_api_key_encrypt_decrypt():
|
||||||
|
"""encrypt_api_key + decrypt_api_key round-trip returns original plaintext.
|
||||||
|
|
||||||
|
Also verifies domain salt isolation: different provider_ids produce different
|
||||||
|
ciphertexts from the same master key and plaintext (T-07-02).
|
||||||
|
"""
|
||||||
|
from services.ai_config import encrypt_api_key, decrypt_api_key
|
||||||
|
|
||||||
|
# Use a 32-byte master key (CLOUD_CREDS_KEY is Base64-encoded in production,
|
||||||
|
# but the helpers accept raw bytes — use ASCII bytes for simplicity here)
|
||||||
|
master_key = b"a" * 32 # 32 bytes deterministic test key
|
||||||
|
|
||||||
|
# Round-trip: encrypt then decrypt must return original plaintext
|
||||||
|
plaintext = "sk-test-api-key-abc123"
|
||||||
|
ciphertext = encrypt_api_key(master_key, "openai", plaintext)
|
||||||
|
recovered = decrypt_api_key(master_key, "openai", ciphertext)
|
||||||
|
assert recovered == plaintext, "Decrypted value must match original plaintext"
|
||||||
|
|
||||||
|
# Domain salt isolation: same key + same plaintext, different provider_id → different ciphertext
|
||||||
|
ciphertext_anthropic = encrypt_api_key(master_key, "anthropic", plaintext)
|
||||||
|
assert ciphertext != ciphertext_anthropic, (
|
||||||
|
"Different provider_ids must produce different ciphertexts "
|
||||||
|
"(HKDF salt isolation — T-07-02)"
|
||||||
|
)
|
||||||
|
|
||||||
|
# Cross-provider decrypt must fail (Fernet raises InvalidToken)
|
||||||
|
from cryptography.fernet import InvalidToken
|
||||||
|
with pytest.raises(InvalidToken):
|
||||||
|
decrypt_api_key(master_key, "anthropic", ciphertext) # openai key, anthropic provider_id
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.xfail(strict=False, reason="Wave 0 stub — promoted in Plan 07-01")
|
# ---------------------------------------------------------------------------
|
||||||
async def test_api_key_encrypt_decrypt():
|
# test_load_provider_config — D-04 DB integration (skipped without PostgreSQL)
|
||||||
pytest.xfail("not implemented yet — Plan 07-01")
|
# ---------------------------------------------------------------------------
|
||||||
|
|
||||||
|
@pytest.mark.skipif(
|
||||||
|
not os.getenv("INTEGRATION"),
|
||||||
|
reason="needs PostgreSQL — set INTEGRATION=1 to run",
|
||||||
|
)
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_load_provider_config(db_session):
|
||||||
|
"""load_provider_config() reads the active row from system_settings and decrypts the api_key.
|
||||||
|
|
||||||
|
Requires a live PostgreSQL session (db_session fixture from conftest.py).
|
||||||
|
Run with: INTEGRATION=1 pytest tests/test_ai_config.py::test_load_provider_config
|
||||||
|
"""
|
||||||
|
pytest.importorskip("psycopg")
|
||||||
|
|
||||||
|
from services.ai_config import encrypt_api_key, load_provider_config
|
||||||
|
from db.models import SystemSettings
|
||||||
|
import uuid
|
||||||
|
|
||||||
|
master_key = b"b" * 32 # deterministic test key
|
||||||
|
test_api_key = "sk-integration-test-key"
|
||||||
|
provider_id = f"test-provider-{uuid.uuid4().hex[:8]}"
|
||||||
|
|
||||||
|
# Encrypt the API key
|
||||||
|
api_key_enc = encrypt_api_key(master_key, provider_id, test_api_key)
|
||||||
|
|
||||||
|
# Insert a SystemSettings row with is_active=True
|
||||||
|
row = SystemSettings(
|
||||||
|
provider_id=provider_id,
|
||||||
|
api_key_enc=api_key_enc,
|
||||||
|
base_url="https://api.example.com/v1",
|
||||||
|
model_name="test-model",
|
||||||
|
context_chars=64000,
|
||||||
|
is_active=True,
|
||||||
|
)
|
||||||
|
db_session.add(row)
|
||||||
|
await db_session.flush()
|
||||||
|
|
||||||
|
# Patch settings.cloud_creds_key to match our test master key
|
||||||
|
from unittest.mock import patch
|
||||||
|
with patch("services.ai_config.settings") as mock_settings:
|
||||||
|
mock_settings.cloud_creds_key = master_key.decode("utf-8")
|
||||||
|
result = await load_provider_config(db_session)
|
||||||
|
|
||||||
|
assert result is not None, "load_provider_config must return a ProviderConfig"
|
||||||
|
assert result.provider_id == provider_id
|
||||||
|
assert result.api_key == test_api_key, "api_key must be decrypted correctly"
|
||||||
|
assert result.base_url == "https://api.example.com/v1"
|
||||||
|
assert result.model == "test-model"
|
||||||
|
assert result.context_chars == 64000
|
||||||
|
|||||||
@@ -5,11 +5,12 @@ Wave 2 (Plan 07-02) promotes: test_get_provider_typed, test_client_singleton,
|
|||||||
test_generic_openai_json_mode, test_context_chars_truncation, test_smart_truncation,
|
test_generic_openai_json_mode, test_context_chars_truncation, test_smart_truncation,
|
||||||
test_gemini_fallback_to_parse_classification.
|
test_gemini_fallback_to_parse_classification.
|
||||||
|
|
||||||
Remaining stubs (promoted in Plan 07-03): test_anthropic_structured_output.
|
Wave 3 (Plan 07-03) promotes: test_anthropic_structured_output.
|
||||||
"""
|
"""
|
||||||
import pytest
|
import pytest
|
||||||
from unittest.mock import AsyncMock, MagicMock, patch
|
from unittest.mock import AsyncMock, MagicMock, patch
|
||||||
|
|
||||||
|
from ai.anthropic_provider import AnthropicProvider, _CLASSIFICATION_SCHEMA
|
||||||
from ai.generic_openai_provider import GenericOpenAIProvider
|
from ai.generic_openai_provider import GenericOpenAIProvider
|
||||||
from ai.openai_provider import OpenAIProvider
|
from ai.openai_provider import OpenAIProvider
|
||||||
from ai.provider_config import ProviderConfig, PROVIDER_DEFAULTS
|
from ai.provider_config import ProviderConfig, PROVIDER_DEFAULTS
|
||||||
@@ -189,9 +190,81 @@ async def test_gemini_fallback_to_parse_classification():
|
|||||||
|
|
||||||
|
|
||||||
# ---------------------------------------------------------------------------
|
# ---------------------------------------------------------------------------
|
||||||
# Stub: promoted in Plan 07-03
|
# Task 1 (Plan 07-03): AnthropicProvider output_config structured output — D-03
|
||||||
# ---------------------------------------------------------------------------
|
# ---------------------------------------------------------------------------
|
||||||
|
|
||||||
@pytest.mark.xfail(strict=False, reason="Wave 0 stub — promoted in Plan 07-03")
|
@pytest.mark.asyncio
|
||||||
async def test_anthropic_structured_output():
|
async def test_anthropic_structured_output():
|
||||||
pytest.xfail("not implemented yet — Plan 07-03")
|
"""AnthropicProvider.classify() passes output_config with the classification schema (D-03).
|
||||||
|
|
||||||
|
Verifies:
|
||||||
|
- output_config kwarg is present in the messages.create call
|
||||||
|
- output_config value matches _CLASSIFICATION_SCHEMA exactly
|
||||||
|
- Provider correctly reads response.content[0].text when stop_reason == "end_turn"
|
||||||
|
- Singleton _client is reused (AsyncAnthropic constructed once per provider instance)
|
||||||
|
"""
|
||||||
|
provider = AnthropicProvider(api_key="test-key", model="claude-sonnet-4-6", context_chars=100)
|
||||||
|
|
||||||
|
# Build a stub response that models a successful end_turn response
|
||||||
|
stub_content = MagicMock()
|
||||||
|
stub_content.text = '{"assigned_topics":[],"new_topic_suggestions":[]}'
|
||||||
|
stub_response = MagicMock()
|
||||||
|
stub_response.content = [stub_content]
|
||||||
|
stub_response.stop_reason = "end_turn"
|
||||||
|
|
||||||
|
mock_create = AsyncMock(return_value=stub_response)
|
||||||
|
|
||||||
|
with patch("ai.anthropic_provider.anthropic.AsyncAnthropic") as mock_cls:
|
||||||
|
mock_client = MagicMock()
|
||||||
|
mock_client.messages = MagicMock()
|
||||||
|
mock_client.messages.create = mock_create
|
||||||
|
mock_cls.return_value = mock_client
|
||||||
|
|
||||||
|
# Re-create provider inside the patch so self._client uses the mock
|
||||||
|
provider = AnthropicProvider(
|
||||||
|
api_key="test-key", model="claude-sonnet-4-6", context_chars=100
|
||||||
|
)
|
||||||
|
result = await provider.classify("short doc text", [], "sys prompt")
|
||||||
|
|
||||||
|
# output_config must be present and match the schema (D-03)
|
||||||
|
call_kwargs = mock_create.await_args.kwargs
|
||||||
|
assert "output_config" in call_kwargs, "output_config must be passed to messages.create()"
|
||||||
|
assert call_kwargs["output_config"] == {
|
||||||
|
"format": {"type": "json_schema", "schema": _CLASSIFICATION_SCHEMA}
|
||||||
|
}, "output_config must use _CLASSIFICATION_SCHEMA"
|
||||||
|
|
||||||
|
# AsyncAnthropic must have been constructed exactly once (D-07 singleton)
|
||||||
|
assert mock_cls.call_count == 1, "AsyncAnthropic must be constructed once (singleton)"
|
||||||
|
|
||||||
|
# Result must be a valid ClassificationResult
|
||||||
|
assert result.topics == []
|
||||||
|
assert result.suggested_new_topics == []
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_anthropic_stop_reason_fallback():
|
||||||
|
"""When stop_reason is not 'end_turn', AnthropicProvider falls back to empty ClassificationResult.
|
||||||
|
|
||||||
|
T-07-08: refusal or max_tokens stop_reason must not crash — parse_classification("") returns
|
||||||
|
an empty result rather than raising an exception.
|
||||||
|
"""
|
||||||
|
stub_content = MagicMock()
|
||||||
|
stub_content.text = '{"assigned_topics":["should","be","ignored"]}'
|
||||||
|
stub_response = MagicMock()
|
||||||
|
stub_response.content = [stub_content]
|
||||||
|
stub_response.stop_reason = "max_tokens" # simulated refusal / truncation
|
||||||
|
|
||||||
|
mock_create = AsyncMock(return_value=stub_response)
|
||||||
|
|
||||||
|
with patch("ai.anthropic_provider.anthropic.AsyncAnthropic") as mock_cls:
|
||||||
|
mock_client = MagicMock()
|
||||||
|
mock_client.messages = MagicMock()
|
||||||
|
mock_client.messages.create = mock_create
|
||||||
|
mock_cls.return_value = mock_client
|
||||||
|
|
||||||
|
provider = AnthropicProvider(api_key="k", model="claude-sonnet-4-6", context_chars=1000)
|
||||||
|
result = await provider.classify("doc text", [], "sys")
|
||||||
|
|
||||||
|
# stop_reason != "end_turn" → raw == "" → empty ClassificationResult, no crash
|
||||||
|
assert result.topics == []
|
||||||
|
assert result.suggested_new_topics == []
|
||||||
|
|||||||
@@ -119,13 +119,17 @@ async def test_classifier_with_mock_provider(isolated_data_dir):
|
|||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
async def test_per_user_provider(db_session):
|
async def test_per_user_provider(db_session):
|
||||||
"""When ai_provider='openai' and ai_model='gpt-4o' are passed to the classifier,
|
"""When ai_provider='openai' and ai_model='gpt-4o' are passed to the classifier,
|
||||||
it resolves _settings['active_provider'] == 'openai'.
|
it builds a ProviderConfig with provider_id='openai' and model='gpt-4o'.
|
||||||
|
|
||||||
DOC-03: AI provider/model comes from the user's DB record (passed through from
|
DOC-03: AI provider/model comes from the user's DB record (passed through from
|
||||||
_run) not from global config or the retired load_settings() flat file (D-14).
|
_run) not from global config or the retired load_settings() flat file (D-14).
|
||||||
|
|
||||||
|
Plan 07-03: get_provider now receives a ProviderConfig object (D-06), not a
|
||||||
|
raw dict. Assertions updated accordingly.
|
||||||
"""
|
"""
|
||||||
from unittest.mock import AsyncMock, patch, MagicMock
|
from unittest.mock import AsyncMock, patch, MagicMock
|
||||||
from ai.base import ClassificationResult
|
from ai.base import ClassificationResult
|
||||||
|
from ai.provider_config import ProviderConfig
|
||||||
from services.classifier import classify_document
|
from services.classifier import classify_document
|
||||||
import uuid
|
import uuid
|
||||||
|
|
||||||
@@ -136,10 +140,10 @@ async def test_per_user_provider(db_session):
|
|||||||
mock_doc = MagicMock()
|
mock_doc = MagicMock()
|
||||||
mock_doc.user_id = user_id
|
mock_doc.user_id = user_id
|
||||||
|
|
||||||
captured_settings = {}
|
captured_configs = []
|
||||||
|
|
||||||
def capture_get_provider(settings):
|
def capture_get_provider(config):
|
||||||
captured_settings.update(settings)
|
captured_configs.append(config)
|
||||||
mock_provider = MagicMock()
|
mock_provider = MagicMock()
|
||||||
mock_provider.classify = AsyncMock(return_value=ClassificationResult(
|
mock_provider.classify = AsyncMock(return_value=ClassificationResult(
|
||||||
topics=[], suggested_new_topics=[], reasoning=""
|
topics=[], suggested_new_topics=[], reasoning=""
|
||||||
@@ -156,9 +160,11 @@ async def test_per_user_provider(db_session):
|
|||||||
patch("services.classifier.get_provider", side_effect=capture_get_provider):
|
patch("services.classifier.get_provider", side_effect=capture_get_provider):
|
||||||
await classify_document(mock_session, doc_id, ai_provider="openai", ai_model="gpt-4o")
|
await classify_document(mock_session, doc_id, ai_provider="openai", ai_model="gpt-4o")
|
||||||
|
|
||||||
assert captured_settings.get("active_provider") == "openai"
|
assert len(captured_configs) == 1
|
||||||
assert "openai" in captured_settings.get("providers", {})
|
config = captured_configs[0]
|
||||||
assert captured_settings["providers"]["openai"]["model"] == "gpt-4o"
|
assert isinstance(config, ProviderConfig), "get_provider must receive a ProviderConfig (D-06)"
|
||||||
|
assert config.provider_id == "openai"
|
||||||
|
assert config.model == "gpt-4o"
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
@@ -224,13 +230,18 @@ async def test_celery_task_uses_user_provider(db_session):
|
|||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
async def test_default_provider_fallback(db_session):
|
async def test_default_provider_fallback(db_session):
|
||||||
"""When user.ai_provider is None, the classifier receives config.settings.default_ai_provider.
|
"""When user.ai_provider is None, the classifier uses the system_settings DB row or env fallback.
|
||||||
|
|
||||||
D-15: fallback chain is user.ai_provider → DEFAULT_AI_PROVIDER env var →
|
D-15: fallback chain is user.ai_provider → system_settings DB → DEFAULT_AI_PROVIDER env var →
|
||||||
code default 'ollama' (CONTEXT.md D-15).
|
code default 'ollama' (CONTEXT.md D-15).
|
||||||
|
|
||||||
|
Plan 07-03: get_provider now receives a ProviderConfig (D-06). When ai_provider=None
|
||||||
|
and load_provider_config returns None, the classifier falls back to
|
||||||
|
app_settings.default_ai_provider. Assertions updated from dict to ProviderConfig.
|
||||||
"""
|
"""
|
||||||
from unittest.mock import AsyncMock, patch, MagicMock
|
from unittest.mock import AsyncMock, patch, MagicMock
|
||||||
from ai.base import ClassificationResult
|
from ai.base import ClassificationResult
|
||||||
|
from ai.provider_config import ProviderConfig
|
||||||
from services.classifier import classify_document
|
from services.classifier import classify_document
|
||||||
import uuid
|
import uuid
|
||||||
|
|
||||||
@@ -241,10 +252,10 @@ async def test_default_provider_fallback(db_session):
|
|||||||
mock_doc = MagicMock()
|
mock_doc = MagicMock()
|
||||||
mock_doc.user_id = user_id
|
mock_doc.user_id = user_id
|
||||||
|
|
||||||
captured_settings = {}
|
captured_configs = []
|
||||||
|
|
||||||
def capture_get_provider(settings):
|
def capture_get_provider(config):
|
||||||
captured_settings.update(settings)
|
captured_configs.append(config)
|
||||||
mock_provider = MagicMock()
|
mock_provider = MagicMock()
|
||||||
mock_provider.classify = AsyncMock(return_value=ClassificationResult(
|
mock_provider.classify = AsyncMock(return_value=ClassificationResult(
|
||||||
topics=[], suggested_new_topics=[], reasoning=""
|
topics=[], suggested_new_topics=[], reasoning=""
|
||||||
@@ -255,6 +266,7 @@ async def test_default_provider_fallback(db_session):
|
|||||||
patch("services.classifier.storage.load_topics_for_user", AsyncMock(return_value=[])), \
|
patch("services.classifier.storage.load_topics_for_user", AsyncMock(return_value=[])), \
|
||||||
patch("services.classifier.storage.load_topics", AsyncMock(return_value=[])), \
|
patch("services.classifier.storage.load_topics", AsyncMock(return_value=[])), \
|
||||||
patch("services.classifier.storage.update_document_topics", AsyncMock(return_value=None)), \
|
patch("services.classifier.storage.update_document_topics", AsyncMock(return_value=None)), \
|
||||||
|
patch("services.classifier.load_provider_config", AsyncMock(return_value=None)), \
|
||||||
patch("services.classifier.get_provider", side_effect=capture_get_provider):
|
patch("services.classifier.get_provider", side_effect=capture_get_provider):
|
||||||
|
|
||||||
mock_session = AsyncMock()
|
mock_session = AsyncMock()
|
||||||
@@ -262,5 +274,8 @@ async def test_default_provider_fallback(db_session):
|
|||||||
# Pass ai_provider=None to trigger the default fallback (D-15)
|
# Pass ai_provider=None to trigger the default fallback (D-15)
|
||||||
await classify_document(mock_session, doc_id, ai_provider=None, ai_model=None)
|
await classify_document(mock_session, doc_id, ai_provider=None, ai_model=None)
|
||||||
|
|
||||||
|
assert len(captured_configs) == 1
|
||||||
|
config = captured_configs[0]
|
||||||
|
assert isinstance(config, ProviderConfig), "get_provider must receive a ProviderConfig (D-06)"
|
||||||
# Should fall back to app_settings.default_ai_provider = "ollama"
|
# Should fall back to app_settings.default_ai_provider = "ollama"
|
||||||
assert captured_settings.get("active_provider") == "ollama"
|
assert config.provider_id == "ollama"
|
||||||
|
|||||||
Reference in New Issue
Block a user