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}
|
||||
|
||||
|
||||
|
||||
Reference in New Issue
Block a user