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:
cesnimda
2026-08-09 12:30:11 +02:00
parent c3f4a57195
commit 5eb9b3cb96
29 changed files with 967 additions and 145 deletions
+225 -33
View File
@@ -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}