feat(ai): enforce local-first routing
Keep external providers behind server consent, task, and prompt-cost gates while persisting actual provider provenance.
This commit is contained in:
+225
-33
@@ -15,6 +15,8 @@ import os
|
||||
import re
|
||||
import torch
|
||||
import pytesseract
|
||||
import threading
|
||||
import time
|
||||
from urllib import request as urllib_request
|
||||
from urllib.error import URLError, HTTPError
|
||||
from contextvars import ContextVar
|
||||
@@ -35,7 +37,18 @@ AI_SERVICE_TOKEN_HEADER = "X-Ai-Service-Token"
|
||||
# exposes no user data and no generation path.
|
||||
AI_SERVICE_OPEN_PATHS = {"/health"}
|
||||
EXTERNAL_AI_ALLOWED_HEADER = "X-Ai-External-Allowed"
|
||||
AI_TASK_TYPE_HEADER = "X-Ai-Task-Type"
|
||||
_external_ai_allowed = ContextVar("external_ai_allowed", default=False)
|
||||
_ai_task_type = ContextVar("ai_task_type", default="unknown")
|
||||
_route_state = ContextVar("route_state", default=None)
|
||||
|
||||
|
||||
_PATH_TASKS = {
|
||||
"/cv/normalize": "cv-normalize",
|
||||
"/cv/classify-block": "cv-classify",
|
||||
"/cv/rewrite": "cv-rewrite",
|
||||
"/summarize": "job-summary",
|
||||
}
|
||||
|
||||
|
||||
@app.middleware("http")
|
||||
@@ -52,14 +65,28 @@ async def require_service_token(request: Request, call_next):
|
||||
EXTERNAL_AI_ENABLED
|
||||
and request.headers.get(EXTERNAL_AI_ALLOWED_HEADER, "").strip().lower() == "true"
|
||||
)
|
||||
token = _external_ai_allowed.set(allowed)
|
||||
requested_task = request.headers.get(AI_TASK_TYPE_HEADER, "").strip().lower()
|
||||
task_type = requested_task if re.fullmatch(r"[a-z0-9._-]{1,64}", requested_task) else _PATH_TASKS.get(request.url.path, "unknown")
|
||||
state = {"provider": None, "model": None, "fallback_reason": None, "route_reason": None}
|
||||
allowed_token = _external_ai_allowed.set(allowed)
|
||||
task_token = _ai_task_type.set(task_type)
|
||||
state_token = _route_state.set(state)
|
||||
try:
|
||||
response = await call_next(request)
|
||||
if request.url.path.startswith("/cv/"):
|
||||
response.headers["X-Ai-Provider"] = _effective_provider()
|
||||
if state["provider"]:
|
||||
response.headers["X-Ai-Provider"] = state["provider"]
|
||||
if state["model"]:
|
||||
response.headers["X-Ai-Model"] = state["model"]
|
||||
if state["fallback_reason"]:
|
||||
response.headers["X-Ai-Fallback-Reason"] = state["fallback_reason"]
|
||||
if state["route_reason"]:
|
||||
response.headers["X-Ai-Route-Reason"] = state["route_reason"]
|
||||
return response
|
||||
finally:
|
||||
_external_ai_allowed.reset(token)
|
||||
_route_state.reset(state_token)
|
||||
_ai_task_type.reset(task_token)
|
||||
_external_ai_allowed.reset(allowed_token)
|
||||
|
||||
MODEL_NAME = "sshleifer/distilbart-cnn-12-6"
|
||||
MAX_INPUT_CHARS = 20000
|
||||
@@ -76,6 +103,17 @@ OLLAMA_MODEL = os.getenv("OLLAMA_MODEL", "")
|
||||
# (distilbart) regardless of this setting.
|
||||
AI_PROVIDER = (os.getenv("AI_PROVIDER", "ollama").strip().lower() or "ollama")
|
||||
EXTERNAL_AI_ENABLED = os.getenv("EXTERNAL_AI_ENABLED", "").strip().lower() in {"1", "true", "yes"}
|
||||
AI_ROUTING_MODE = (os.getenv("AI_ROUTING_MODE", "local_first").strip().lower() or "local_first")
|
||||
if AI_ROUTING_MODE not in {"local_only", "local_first", "external_only"}:
|
||||
AI_ROUTING_MODE = "local_only"
|
||||
EXTERNAL_AI_ALLOWED_TASKS = frozenset(
|
||||
item.strip().lower()
|
||||
for item in os.getenv("EXTERNAL_AI_ALLOWED_TASKS", "cv-normalize,cv-classify,cv-rewrite").split(",")
|
||||
if item.strip()
|
||||
)
|
||||
EXTERNAL_AI_MAX_PROMPT_CHARS = max(1000, min(int(os.getenv("EXTERNAL_AI_MAX_PROMPT_CHARS", "24000")), 100000))
|
||||
LOCAL_CIRCUIT_FAILURE_THRESHOLD = max(1, min(int(os.getenv("LOCAL_AI_CIRCUIT_FAILURE_THRESHOLD", "3")), 20))
|
||||
LOCAL_CIRCUIT_OPEN_SECONDS = max(1, min(int(os.getenv("LOCAL_AI_CIRCUIT_OPEN_SECONDS", "30")), 600))
|
||||
GEMINI_API_KEY = os.getenv("GEMINI_API_KEY", "").strip()
|
||||
GEMINI_MODEL = os.getenv("GEMINI_MODEL", "gemini-2.0-flash").strip()
|
||||
GEMINI_BASE_URL = os.getenv("GEMINI_BASE_URL", "https://generativelanguage.googleapis.com").rstrip("/")
|
||||
@@ -85,6 +123,10 @@ GROQ_BASE_URL = os.getenv("GROQ_BASE_URL", "https://api.groq.com/openai/v1").rst
|
||||
SKIP_MODEL_LOAD = os.getenv("AI_SERVICE_SKIP_MODEL_LOAD", "") == "1"
|
||||
EAGER_MODEL_LOAD = os.getenv("AI_SERVICE_EAGER_MODEL_LOAD", "") == "1"
|
||||
|
||||
_local_circuit_lock = threading.Lock()
|
||||
_local_failure_count = 0
|
||||
_local_circuit_open_until = 0.0
|
||||
|
||||
|
||||
tokenizer = None
|
||||
model = None
|
||||
@@ -231,7 +273,12 @@ async def health():
|
||||
"summarize_available": MODEL_LOADED and not MODEL_DISABLED,
|
||||
"model_load_error": MODEL_LOAD_ERROR,
|
||||
"ai_provider": AI_PROVIDER,
|
||||
"ai_provider_configured": _provider_configured(),
|
||||
"ai_provider_configured": _provider_configured(AI_PROVIDER),
|
||||
"ai_routing_mode": AI_ROUTING_MODE,
|
||||
"external_ai_enabled": EXTERNAL_AI_ENABLED,
|
||||
"external_ai_allowed_tasks": sorted(EXTERNAL_AI_ALLOWED_TASKS),
|
||||
"external_ai_max_prompt_chars": EXTERNAL_AI_MAX_PROMPT_CHARS,
|
||||
**_local_circuit_status(),
|
||||
**_ollama_status(),
|
||||
}
|
||||
|
||||
@@ -455,18 +502,94 @@ def _provider_display(provider: str) -> str:
|
||||
return _PROVIDER_DISPLAY.get(provider, provider or "AI provider")
|
||||
|
||||
|
||||
def _provider_configured() -> bool:
|
||||
if AI_PROVIDER == "gemini":
|
||||
def _provider_configured(provider: str | None = None) -> bool:
|
||||
provider = provider or AI_PROVIDER
|
||||
if provider == "gemini":
|
||||
return bool(GEMINI_API_KEY)
|
||||
if AI_PROVIDER == "groq":
|
||||
if provider == "groq":
|
||||
return bool(GROQ_API_KEY)
|
||||
return bool(OLLAMA_MODEL)
|
||||
|
||||
|
||||
def _effective_provider() -> str:
|
||||
if EXTERNAL_AI_ENABLED and _external_ai_allowed.get() and AI_PROVIDER in {"gemini", "groq"}:
|
||||
return AI_PROVIDER
|
||||
return "ollama"
|
||||
def _provider_model(provider: str) -> str | None:
|
||||
return {
|
||||
"ollama": OLLAMA_MODEL,
|
||||
"gemini": GEMINI_MODEL,
|
||||
"groq": GROQ_MODEL,
|
||||
}.get(provider) or None
|
||||
|
||||
|
||||
def _set_route_metadata(provider: str | None, route_reason: str, fallback_reason: str | None = None):
|
||||
state = _route_state.get()
|
||||
if state is None:
|
||||
state = {"provider": None, "model": None, "fallback_reason": None, "route_reason": None}
|
||||
_route_state.set(state)
|
||||
state["provider"] = provider
|
||||
state["model"] = _provider_model(provider) if provider else None
|
||||
state["fallback_reason"] = fallback_reason
|
||||
state["route_reason"] = route_reason
|
||||
|
||||
|
||||
def _local_circuit_is_open() -> bool:
|
||||
global _local_failure_count, _local_circuit_open_until
|
||||
now = time.monotonic()
|
||||
with _local_circuit_lock:
|
||||
if _local_circuit_open_until <= now:
|
||||
_local_circuit_open_until = 0.0
|
||||
if _local_failure_count >= LOCAL_CIRCUIT_FAILURE_THRESHOLD:
|
||||
_local_failure_count = 0
|
||||
return False
|
||||
return True
|
||||
|
||||
|
||||
def _record_local_success():
|
||||
global _local_failure_count, _local_circuit_open_until
|
||||
with _local_circuit_lock:
|
||||
_local_failure_count = 0
|
||||
_local_circuit_open_until = 0.0
|
||||
|
||||
|
||||
def _record_local_failure():
|
||||
global _local_failure_count, _local_circuit_open_until
|
||||
with _local_circuit_lock:
|
||||
_local_failure_count += 1
|
||||
if _local_failure_count >= LOCAL_CIRCUIT_FAILURE_THRESHOLD:
|
||||
_local_circuit_open_until = time.monotonic() + LOCAL_CIRCUIT_OPEN_SECONDS
|
||||
|
||||
|
||||
def _local_circuit_status() -> dict:
|
||||
now = time.monotonic()
|
||||
with _local_circuit_lock:
|
||||
remaining = max(0.0, _local_circuit_open_until - now)
|
||||
return {
|
||||
"local_circuit_open": remaining > 0,
|
||||
"local_circuit_failures": _local_failure_count,
|
||||
"local_circuit_retry_after_seconds": round(remaining, 1),
|
||||
}
|
||||
|
||||
|
||||
class _ProviderFailure(Exception):
|
||||
def __init__(self, provider: str, category: str, status_code: int):
|
||||
super().__init__(category)
|
||||
self.provider = provider
|
||||
self.category = category
|
||||
self.status_code = status_code
|
||||
|
||||
|
||||
def _external_denial_reason(prompt: str) -> str | None:
|
||||
if AI_ROUTING_MODE == "local_only":
|
||||
return "local_only"
|
||||
if not EXTERNAL_AI_ENABLED or not _external_ai_allowed.get():
|
||||
return "external_not_permitted"
|
||||
if AI_PROVIDER not in {"gemini", "groq"}:
|
||||
return "external_not_configured"
|
||||
if _ai_task_type.get() not in EXTERNAL_AI_ALLOWED_TASKS:
|
||||
return "task_not_allowed_external"
|
||||
if not _provider_configured(AI_PROVIDER):
|
||||
return "external_not_configured"
|
||||
if len(prompt) > EXTERNAL_AI_MAX_PROMPT_CHARS:
|
||||
return "external_prompt_limit"
|
||||
return None
|
||||
|
||||
|
||||
def _http_post_json(url: str, payload: dict, headers: dict, timeout: int) -> dict:
|
||||
@@ -534,43 +657,115 @@ def _groq_generate(prompt: str, *, json_mode: bool, temperature: float, timeout:
|
||||
return ((choices[0].get("message") or {}).get("content") or "").strip()
|
||||
|
||||
|
||||
def _provider_generate(prompt: str, *, json_mode: bool, temperature: float, timeout: int) -> str:
|
||||
provider = _effective_provider()
|
||||
def _generate_from_provider(provider: str, prompt: str, *, json_mode: bool, temperature: float, timeout: int) -> str:
|
||||
try:
|
||||
if provider == "gemini":
|
||||
return _gemini_generate(prompt, json_mode=json_mode, temperature=temperature, timeout=timeout)
|
||||
if provider == "groq":
|
||||
return _groq_generate(prompt, json_mode=json_mode, temperature=temperature, timeout=timeout)
|
||||
return _ollama_generate(prompt, json_mode=json_mode, temperature=temperature, timeout=timeout)
|
||||
except HTTPException:
|
||||
raise
|
||||
except HTTPException as ex:
|
||||
category = "provider_not_configured" if ex.status_code == 503 else "provider_rejected"
|
||||
raise _ProviderFailure(provider, category, ex.status_code) from ex
|
||||
except HTTPError as ex:
|
||||
raise HTTPException(status_code=502, detail=f"{_provider_display(provider)} request failed with {ex.code}.")
|
||||
except URLError as ex:
|
||||
raise HTTPException(status_code=503, detail=f"{_provider_display(provider)} is unreachable: {ex.reason}.")
|
||||
category = "provider_busy" if ex.code == 429 else "provider_unavailable"
|
||||
raise _ProviderFailure(provider, category, 503 if ex.code in {408, 429, 502, 503, 504} else 502) from ex
|
||||
except (URLError, TimeoutError) as ex:
|
||||
raise _ProviderFailure(provider, "provider_unavailable", 503) from ex
|
||||
|
||||
|
||||
def _ollama_generate_json(prompt: str):
|
||||
provider = _effective_provider()
|
||||
raw = _provider_generate(prompt, json_mode=True, temperature=0.1, timeout=120)
|
||||
def _parse_provider_json(raw: str, provider: str):
|
||||
if not raw:
|
||||
raise HTTPException(status_code=502, detail=f"{_provider_display(provider)} returned an empty response.")
|
||||
raise _ProviderFailure(provider, "empty_response", 502)
|
||||
try:
|
||||
return json.loads(raw)
|
||||
except json.JSONDecodeError:
|
||||
except json.JSONDecodeError as first_error:
|
||||
start = raw.find("{")
|
||||
end = raw.rfind("}")
|
||||
if start >= 0 and end > start:
|
||||
return json.loads(raw[start:end + 1])
|
||||
raise HTTPException(status_code=502, detail=f"{_provider_display(provider)} did not return valid JSON.")
|
||||
try:
|
||||
return json.loads(raw[start:end + 1])
|
||||
except json.JSONDecodeError:
|
||||
pass
|
||||
raise _ProviderFailure(provider, "schema_invalid", 502) from first_error
|
||||
|
||||
|
||||
def _validated_generation(provider: str, prompt: str, *, json_mode: bool, temperature: float, timeout: int):
|
||||
raw = _generate_from_provider(provider, prompt, json_mode=json_mode, temperature=temperature, timeout=timeout)
|
||||
if json_mode:
|
||||
return _parse_provider_json(raw, provider)
|
||||
if not raw:
|
||||
raise _ProviderFailure(provider, "empty_response", 502)
|
||||
return raw
|
||||
|
||||
|
||||
def _raise_route_failure(failure: _ProviderFailure):
|
||||
display = _provider_display(failure.provider)
|
||||
messages = {
|
||||
"provider_not_configured": f"{display} is not configured.",
|
||||
"provider_busy": f"{display} is busy. Try again later.",
|
||||
"schema_invalid": f"{display} returned an invalid structured response.",
|
||||
"empty_response": f"{display} returned an empty response.",
|
||||
"provider_rejected": f"{display} rejected the request.",
|
||||
}
|
||||
raise HTTPException(
|
||||
status_code=failure.status_code,
|
||||
detail=messages.get(failure.category, f"{display} is unavailable."),
|
||||
)
|
||||
|
||||
|
||||
def _route_generation(prompt: str, *, json_mode: bool, temperature: float, timeout: int):
|
||||
denial_reason = _external_denial_reason(prompt)
|
||||
|
||||
if AI_ROUTING_MODE == "external_only":
|
||||
if denial_reason is not None:
|
||||
_set_route_metadata(None, denial_reason)
|
||||
raise HTTPException(status_code=403, detail="External AI processing is not permitted for this request.")
|
||||
try:
|
||||
result = _validated_generation(AI_PROVIDER, prompt, json_mode=json_mode, temperature=temperature, timeout=timeout)
|
||||
_set_route_metadata(AI_PROVIDER, "external_only")
|
||||
return result
|
||||
except _ProviderFailure as failure:
|
||||
_set_route_metadata(failure.provider, failure.category)
|
||||
_raise_route_failure(failure)
|
||||
|
||||
if _local_circuit_is_open():
|
||||
if denial_reason is None:
|
||||
try:
|
||||
result = _validated_generation(AI_PROVIDER, prompt, json_mode=json_mode, temperature=temperature, timeout=timeout)
|
||||
_set_route_metadata(AI_PROVIDER, "external_fallback", "local_circuit_open")
|
||||
return result
|
||||
except _ProviderFailure as failure:
|
||||
_set_route_metadata(failure.provider, failure.category, "local_circuit_open")
|
||||
_raise_route_failure(failure)
|
||||
_set_route_metadata(None, f"local_circuit_open:{denial_reason}")
|
||||
raise HTTPException(status_code=503, detail="Local AI is temporarily unavailable. Try again later.")
|
||||
|
||||
try:
|
||||
result = _validated_generation("ollama", prompt, json_mode=json_mode, temperature=temperature, timeout=timeout)
|
||||
_record_local_success()
|
||||
_set_route_metadata("ollama", "local_primary")
|
||||
return result
|
||||
except _ProviderFailure as local_failure:
|
||||
_record_local_failure()
|
||||
if denial_reason is not None:
|
||||
_set_route_metadata("ollama", f"{local_failure.category}:{denial_reason}")
|
||||
_raise_route_failure(local_failure)
|
||||
try:
|
||||
result = _validated_generation(AI_PROVIDER, prompt, json_mode=json_mode, temperature=temperature, timeout=timeout)
|
||||
_set_route_metadata(AI_PROVIDER, "external_fallback", f"local_{local_failure.category}")
|
||||
return result
|
||||
except _ProviderFailure as external_failure:
|
||||
_set_route_metadata(external_failure.provider, external_failure.category, f"local_{local_failure.category}")
|
||||
_raise_route_failure(external_failure)
|
||||
|
||||
|
||||
def _ollama_generate_json(prompt: str):
|
||||
return _route_generation(prompt, json_mode=True, temperature=0.1, timeout=120)
|
||||
|
||||
|
||||
def _ollama_generate_text(prompt: str) -> str:
|
||||
provider = _effective_provider()
|
||||
raw = _provider_generate(prompt, json_mode=False, temperature=0.2, timeout=180)
|
||||
if not raw:
|
||||
raise HTTPException(status_code=502, detail=f"{_provider_display(provider)} returned an empty rewrite.")
|
||||
return raw
|
||||
return _route_generation(prompt, json_mode=False, temperature=0.2, timeout=180)
|
||||
|
||||
|
||||
@app.post("/cv/normalize")
|
||||
@@ -755,9 +950,6 @@ section.
|
||||
""".strip()
|
||||
|
||||
rewritten = _ollama_generate_text(prompt).strip()
|
||||
if not rewritten:
|
||||
raise HTTPException(status_code=502, detail="Ollama returned an empty rewrite.")
|
||||
|
||||
return {"rewritten_text": rewritten}
|
||||
|
||||
|
||||
|
||||
@@ -11,7 +11,17 @@ if str(ROOT) not in sys.path:
|
||||
sys.path.insert(0, str(ROOT))
|
||||
|
||||
|
||||
def load_app_module(monkeypatch, *, skip_model_load=True, ollama_model=None, service_token=None, external_ai_enabled=False):
|
||||
def load_app_module(
|
||||
monkeypatch,
|
||||
*,
|
||||
skip_model_load=True,
|
||||
ollama_model=None,
|
||||
service_token=None,
|
||||
external_ai_enabled=False,
|
||||
routing_mode="local_first",
|
||||
circuit_threshold=3,
|
||||
external_prompt_limit=24000,
|
||||
):
|
||||
if skip_model_load:
|
||||
monkeypatch.setenv("AI_SERVICE_SKIP_MODEL_LOAD", "1")
|
||||
else:
|
||||
@@ -30,6 +40,10 @@ def load_app_module(monkeypatch, *, skip_model_load=True, ollama_model=None, ser
|
||||
monkeypatch.setenv("EXTERNAL_AI_ENABLED", "true")
|
||||
else:
|
||||
monkeypatch.delenv("EXTERNAL_AI_ENABLED", raising=False)
|
||||
monkeypatch.setenv("AI_ROUTING_MODE", routing_mode)
|
||||
monkeypatch.setenv("LOCAL_AI_CIRCUIT_FAILURE_THRESHOLD", str(circuit_threshold))
|
||||
monkeypatch.setenv("EXTERNAL_AI_MAX_PROMPT_CHARS", str(external_prompt_limit))
|
||||
monkeypatch.delenv("EXTERNAL_AI_ALLOWED_TASKS", raising=False)
|
||||
if "app" in sys.modules:
|
||||
del sys.modules["app"]
|
||||
module = importlib.import_module("app")
|
||||
@@ -227,56 +241,62 @@ def test_provider_defaults_to_ollama_and_is_unchanged(monkeypatch):
|
||||
assert captured["body"]["options"]["temperature"] == 0.1
|
||||
|
||||
|
||||
def test_provider_gemini_dispatch(monkeypatch):
|
||||
def test_local_success_wins_even_when_external_fallback_is_permitted(monkeypatch):
|
||||
monkeypatch.setenv("AI_PROVIDER", "gemini")
|
||||
monkeypatch.setenv("GEMINI_API_KEY", "test-key")
|
||||
monkeypatch.setenv("GEMINI_MODEL", "gemini-2.0-flash")
|
||||
module = load_app_module(monkeypatch, external_ai_enabled=True)
|
||||
module = load_app_module(monkeypatch, ollama_model="qwen2.5:7b", external_ai_enabled=True)
|
||||
module._external_ai_allowed.set(True)
|
||||
module._ai_task_type.set("cv-normalize")
|
||||
|
||||
captured = {}
|
||||
payload = {"candidates": [{"content": {"parts": [{"text": '{"score": 9}'}]}}]}
|
||||
_install_fake_urlopen(monkeypatch, module, payload, captured)
|
||||
calls = []
|
||||
monkeypatch.setattr(module, "_ollama_generate", lambda *args, **kwargs: calls.append("ollama") or '{"score": 7}')
|
||||
monkeypatch.setattr(module, "_gemini_generate", lambda *args, **kwargs: calls.append("gemini") or '{"score": 9}')
|
||||
|
||||
assert module._ollama_generate_json("hi") == {"score": 9}
|
||||
assert "generativelanguage" in captured["url"]
|
||||
assert "gemini-2.0-flash:generateContent" in captured["url"]
|
||||
assert "key=" not in captured["url"] # key must not be in the URL
|
||||
assert captured["headers"].get("x-goog-api-key") == "test-key"
|
||||
assert captured["body"]["generationConfig"]["responseMimeType"] == "application/json"
|
||||
assert module._ollama_generate_json("hi") == {"score": 7}
|
||||
assert calls == ["ollama"]
|
||||
|
||||
|
||||
def test_provider_groq_dispatch(monkeypatch):
|
||||
def test_local_failure_uses_permitted_groq_fallback_sequentially(monkeypatch):
|
||||
monkeypatch.setenv("AI_PROVIDER", "groq")
|
||||
monkeypatch.setenv("GROQ_API_KEY", "test-key")
|
||||
module = load_app_module(monkeypatch, external_ai_enabled=True)
|
||||
module = load_app_module(monkeypatch, ollama_model="qwen2.5:7b", external_ai_enabled=True)
|
||||
module._external_ai_allowed.set(True)
|
||||
module._ai_task_type.set("cv-rewrite")
|
||||
|
||||
captured = {}
|
||||
payload = {"choices": [{"message": {"content": "rewritten CV text"}}]}
|
||||
_install_fake_urlopen(monkeypatch, module, payload, captured)
|
||||
calls = []
|
||||
|
||||
def local_failure(*args, **kwargs):
|
||||
calls.append("ollama")
|
||||
raise module.URLError("synthetic local outage")
|
||||
|
||||
monkeypatch.setattr(module, "_ollama_generate", local_failure)
|
||||
monkeypatch.setattr(module, "_groq_generate", lambda *args, **kwargs: calls.append("groq") or "rewritten CV text")
|
||||
|
||||
assert module._ollama_generate_text("rewrite this") == "rewritten CV text"
|
||||
assert captured["url"].endswith("/chat/completions")
|
||||
assert captured["headers"].get("authorization") == "Bearer test-key"
|
||||
assert captured["body"]["messages"][0]["content"] == "rewrite this"
|
||||
assert calls == ["ollama", "groq"]
|
||||
assert module._route_state.get()["provider"] == "groq"
|
||||
assert module._route_state.get()["fallback_reason"] == "local_provider_unavailable"
|
||||
|
||||
|
||||
def test_provider_missing_cloud_key_raises_503(monkeypatch):
|
||||
def test_missing_external_key_never_bypasses_local_failure(monkeypatch):
|
||||
monkeypatch.setenv("AI_PROVIDER", "gemini")
|
||||
monkeypatch.delenv("GEMINI_API_KEY", raising=False)
|
||||
module = load_app_module(monkeypatch, external_ai_enabled=True)
|
||||
module = load_app_module(monkeypatch, ollama_model="qwen2.5:7b", external_ai_enabled=True)
|
||||
module._external_ai_allowed.set(True)
|
||||
module._ai_task_type.set("cv-normalize")
|
||||
monkeypatch.setattr(module, "_ollama_generate", lambda *args, **kwargs: (_ for _ in ()).throw(module.URLError("synthetic outage")))
|
||||
monkeypatch.setattr(module, "_gemini_generate", lambda *args, **kwargs: (_ for _ in ()).throw(AssertionError("must not call Gemini")))
|
||||
|
||||
from fastapi import HTTPException
|
||||
|
||||
try:
|
||||
module._ollama_generate_json("hi")
|
||||
except HTTPException as ex:
|
||||
except module.HTTPException as ex:
|
||||
assert ex.status_code == 503
|
||||
assert "GEMINI_API_KEY" in ex.detail
|
||||
assert "Ollama" in ex.detail
|
||||
else:
|
||||
raise AssertionError("expected HTTPException for missing GEMINI_API_KEY")
|
||||
raise AssertionError("expected the local failure")
|
||||
|
||||
|
||||
def test_health_reports_active_provider(monkeypatch):
|
||||
@@ -289,31 +309,161 @@ def test_health_reports_active_provider(monkeypatch):
|
||||
|
||||
assert payload["ai_provider"] == "gemini"
|
||||
assert payload["ai_provider_configured"] is True
|
||||
assert payload["ai_routing_mode"] == "local_first"
|
||||
assert payload["local_circuit_open"] is False
|
||||
|
||||
|
||||
def test_external_provider_requires_admin_gate_and_backend_consent_header(monkeypatch):
|
||||
def test_external_fallback_requires_admin_gate_and_backend_consent_header(monkeypatch):
|
||||
monkeypatch.setenv("AI_PROVIDER", "gemini")
|
||||
monkeypatch.setenv("GEMINI_API_KEY", "test-key")
|
||||
module = load_app_module(monkeypatch, ollama_model="qwen2.5:7b", external_ai_enabled=True)
|
||||
captured = {}
|
||||
_install_fake_urlopen(monkeypatch, module, {"response": "rewritten locally"}, captured)
|
||||
calls = []
|
||||
|
||||
def fake_generate(provider, *args, **kwargs):
|
||||
calls.append(provider)
|
||||
if provider == "ollama":
|
||||
raise module._ProviderFailure("ollama", "provider_unavailable", 503)
|
||||
return "rewritten externally"
|
||||
|
||||
monkeypatch.setattr(module, "_generate_from_provider", fake_generate)
|
||||
client = TestClient(module.app)
|
||||
|
||||
local_response = client.post("/cv/rewrite", json={"instruction": "Rewrite", "text": "Synthetic CV"})
|
||||
assert local_response.status_code == 200
|
||||
assert captured["url"].startswith("http://127.0.0.1:11434/")
|
||||
assert local_response.status_code == 503
|
||||
assert calls == ["ollama"]
|
||||
assert local_response.headers["X-Ai-Provider"] == "ollama"
|
||||
|
||||
external_payload = {"candidates": [{"content": {"parts": [{"text": "rewritten externally"}]}}]}
|
||||
_install_fake_urlopen(monkeypatch, module, external_payload, captured)
|
||||
calls.clear()
|
||||
external_response = client.post(
|
||||
"/cv/rewrite",
|
||||
json={"instruction": "Rewrite", "text": "Synthetic CV"},
|
||||
headers={"X-Ai-External-Allowed": "true"},
|
||||
)
|
||||
assert external_response.status_code == 200
|
||||
assert "generativelanguage" in captured["url"]
|
||||
assert calls == ["ollama", "gemini"]
|
||||
assert external_response.headers["X-Ai-Provider"] == "gemini"
|
||||
assert external_response.headers["X-Ai-Model"] == "gemini-2.0-flash"
|
||||
assert external_response.headers["X-Ai-Fallback-Reason"] == "local_provider_unavailable"
|
||||
assert external_response.headers["X-Ai-Route-Reason"] == "external_fallback"
|
||||
|
||||
|
||||
def test_invalid_local_json_can_fallback_but_prompt_cost_cap_cannot(monkeypatch):
|
||||
monkeypatch.setenv("AI_PROVIDER", "gemini")
|
||||
monkeypatch.setenv("GEMINI_API_KEY", "test-key")
|
||||
module = load_app_module(
|
||||
monkeypatch,
|
||||
ollama_model="qwen2.5:7b",
|
||||
external_ai_enabled=True,
|
||||
external_prompt_limit=1000,
|
||||
)
|
||||
module._external_ai_allowed.set(True)
|
||||
module._ai_task_type.set("cv-normalize")
|
||||
calls = []
|
||||
monkeypatch.setattr(module, "_ollama_generate", lambda *args, **kwargs: calls.append("ollama") or "not json")
|
||||
monkeypatch.setattr(module, "_gemini_generate", lambda *args, **kwargs: calls.append("gemini") or '{"score": 9}')
|
||||
|
||||
assert module._ollama_generate_json("short") == {"score": 9}
|
||||
assert calls == ["ollama", "gemini"]
|
||||
assert module._route_state.get()["fallback_reason"] == "local_schema_invalid"
|
||||
|
||||
calls.clear()
|
||||
try:
|
||||
module._ollama_generate_json("x" * 1001)
|
||||
except module.HTTPException as ex:
|
||||
assert ex.status_code == 502
|
||||
else:
|
||||
raise AssertionError("expected local schema failure above the external prompt cap")
|
||||
assert calls == ["ollama"]
|
||||
assert "external_prompt_limit" in module._route_state.get()["route_reason"]
|
||||
|
||||
|
||||
def test_open_local_circuit_skips_local_only_when_fallback_is_permitted(monkeypatch):
|
||||
monkeypatch.setenv("AI_PROVIDER", "groq")
|
||||
monkeypatch.setenv("GROQ_API_KEY", "test-key")
|
||||
module = load_app_module(
|
||||
monkeypatch,
|
||||
ollama_model="qwen2.5:7b",
|
||||
external_ai_enabled=True,
|
||||
circuit_threshold=1,
|
||||
)
|
||||
module._external_ai_allowed.set(True)
|
||||
module._ai_task_type.set("cv-rewrite")
|
||||
calls = []
|
||||
|
||||
def local_failure(*args, **kwargs):
|
||||
calls.append("ollama")
|
||||
raise module.URLError("synthetic local outage")
|
||||
|
||||
monkeypatch.setattr(module, "_ollama_generate", local_failure)
|
||||
monkeypatch.setattr(module, "_groq_generate", lambda *args, **kwargs: calls.append("groq") or "external")
|
||||
|
||||
assert module._ollama_generate_text("first") == "external"
|
||||
assert module._ollama_generate_text("second") == "external"
|
||||
assert calls == ["ollama", "groq", "groq"]
|
||||
assert module._route_state.get()["fallback_reason"] == "local_circuit_open"
|
||||
|
||||
|
||||
def test_external_failure_is_clear_and_unapproved_background_task_stays_local(monkeypatch):
|
||||
monkeypatch.setenv("AI_PROVIDER", "gemini")
|
||||
monkeypatch.setenv("GEMINI_API_KEY", "test-key")
|
||||
module = load_app_module(monkeypatch, ollama_model="qwen2.5:7b", external_ai_enabled=True)
|
||||
module._external_ai_allowed.set(True)
|
||||
calls = []
|
||||
|
||||
def provider_failure(provider, *args, **kwargs):
|
||||
calls.append(provider)
|
||||
raise module._ProviderFailure(provider, "provider_unavailable", 503)
|
||||
|
||||
monkeypatch.setattr(module, "_generate_from_provider", provider_failure)
|
||||
module._ai_task_type.set("cv-rewrite")
|
||||
try:
|
||||
module._ollama_generate_text("allowed")
|
||||
except module.HTTPException as ex:
|
||||
assert ex.status_code == 503
|
||||
assert "Gemini" in ex.detail
|
||||
else:
|
||||
raise AssertionError("expected external provider failure")
|
||||
assert calls == ["ollama", "gemini"]
|
||||
assert module._route_state.get()["fallback_reason"] == "local_provider_unavailable"
|
||||
|
||||
calls.clear()
|
||||
module._ai_task_type.set("strategy.snapshot")
|
||||
try:
|
||||
module._ollama_generate_text("not task-approved")
|
||||
except module.HTTPException as ex:
|
||||
assert ex.status_code == 503
|
||||
assert "Ollama" in ex.detail
|
||||
else:
|
||||
raise AssertionError("expected local provider failure")
|
||||
assert calls == ["ollama"]
|
||||
assert "task_not_allowed_external" in module._route_state.get()["route_reason"]
|
||||
|
||||
|
||||
def test_external_only_mode_still_requires_explicit_permission_and_task_allowlist(monkeypatch):
|
||||
monkeypatch.setenv("AI_PROVIDER", "groq")
|
||||
monkeypatch.setenv("GROQ_API_KEY", "test-key")
|
||||
module = load_app_module(
|
||||
monkeypatch,
|
||||
ollama_model="qwen2.5:7b",
|
||||
external_ai_enabled=True,
|
||||
routing_mode="external_only",
|
||||
)
|
||||
module._ai_task_type.set("cv-rewrite")
|
||||
calls = []
|
||||
monkeypatch.setattr(module, "_ollama_generate", lambda *args, **kwargs: calls.append("ollama") or "local")
|
||||
monkeypatch.setattr(module, "_groq_generate", lambda *args, **kwargs: calls.append("groq") or "external")
|
||||
|
||||
try:
|
||||
module._ollama_generate_text("without consent")
|
||||
except module.HTTPException as ex:
|
||||
assert ex.status_code == 403
|
||||
else:
|
||||
raise AssertionError("expected external permission denial")
|
||||
assert calls == []
|
||||
|
||||
module._external_ai_allowed.set(True)
|
||||
assert module._ollama_generate_text("with consent") == "external"
|
||||
assert calls == ["groq"]
|
||||
|
||||
|
||||
def test_service_token_rejects_calls_without_the_header(monkeypatch):
|
||||
|
||||
Reference in New Issue
Block a user