diff --git a/app/ai-service/main.py b/app/ai-service/main.py index 96fccf27..cac8efd2 100644 --- a/app/ai-service/main.py +++ b/app/ai-service/main.py @@ -358,10 +358,15 @@ def settings_customise_sources( ] +from tracing.otel_setup import setup_tracing + + @asynccontextmanager async def lifespan(app: FastAPI): logger.info("Starting up ChainForge AI Service...") + setup_tracing() if not settings.validate_api_keys(): + logger.warning("No API keys configured. AI features will be unavailable.") else: provider = settings.get_active_provider() diff --git a/app/ai-service/requirements-prod.txt b/app/ai-service/requirements-prod.txt index 52046624..d9867753 100644 --- a/app/ai-service/requirements-prod.txt +++ b/app/ai-service/requirements-prod.txt @@ -49,3 +49,7 @@ slowapi==0.1.9 # System monitoring psutil==5.9.8 prometheus-client==0.20.0 +opentelemetry-api==1.25.0 +opentelemetry-sdk==1.25.0 +opentelemetry-exporter-otlp-proto-http==1.25.0 + diff --git a/app/ai-service/requirements.txt b/app/ai-service/requirements.txt index 2f95be39..b28b9470 100644 --- a/app/ai-service/requirements.txt +++ b/app/ai-service/requirements.txt @@ -51,3 +51,7 @@ pytest-cov==4.1.0 # Monitoring and Metrics prometheus-client==0.20.0 psutil==5.9.8 +opentelemetry-api==1.25.0 +opentelemetry-sdk==1.25.0 +opentelemetry-exporter-otlp-proto-http==1.25.0 + diff --git a/app/ai-service/services/humanitarian_verification.py b/app/ai-service/services/humanitarian_verification.py index 65eede6f..d8164ae6 100644 --- a/app/ai-service/services/humanitarian_verification.py +++ b/app/ai-service/services/humanitarian_verification.py @@ -1,343 +1,343 @@ -"""Humanitarian claim verification service with model/provider fallbacks.""" - -import json -import logging -from typing import Any, Dict, List, Optional -import time -import metrics - -import httpx - -from config import settings -from services.humanitarian_prompt import HumanitarianPromptEngine -from services.circuit_breaker import CircuitBreaker -from services.test_provider import TestProvider -from exceptions import AIServiceError - -logger = logging.getLogger(__name__) - - -class HumanitarianVerificationService: - """Runs humanitarian verification against configured LLM providers.""" - - def __init__(self): - self.prompt_engine = HumanitarianPromptEngine() - self.test_provider = TestProvider() - self.breakers = { - "openai": CircuitBreaker( - name="openai", - failure_threshold=settings.circuit_breaker_failure_threshold, - recovery_timeout=settings.circuit_breaker_recovery_timeout_seconds, - ), - "groq": CircuitBreaker( - name="groq", - failure_threshold=settings.circuit_breaker_failure_threshold, - recovery_timeout=settings.circuit_breaker_recovery_timeout_seconds, - ), - } - - def verify_claim( - self, - aid_claim: str, - supporting_evidence: Optional[List[str]] = None, - context_factors: Optional[Dict[str, Any]] = None, - provider_preference: str = "auto", - timeout: Optional[float] = None, - ) -> Dict[str, Any]: - start_time = time.time() - try: - evidence = supporting_evidence or [] - context = context_factors or {} - - primary_prompt = self.prompt_engine.build_primary_prompt( - aid_claim=aid_claim, - supporting_evidence=evidence, - context_factors=context, - ) - fallback_prompt = self.prompt_engine.build_fallback_prompt( - aid_claim=aid_claim, - supporting_evidence=evidence, - context_factors=context, - ) - - providers = self._provider_attempt_order(provider_preference) - if not providers: - raise RuntimeError("No LLM providers configured for humanitarian verification") - - errors: List[str] = [] - - for provider in providers: - breaker = self.breakers.get(provider) - if breaker and not breaker.allow_request(): - logger.warning("Circuit breaker is OPEN for provider=%s. Skipping.", provider) - errors.append(f"provider={provider}, error=Circuit breaker is OPEN") - continue - - model = self._get_model_for_provider(provider) - for prompt_variant, prompt in (("primary", primary_prompt), ("fallback", fallback_prompt)): - try: - logger.info( - "Attempting humanitarian verification with provider=%s model=%s prompt=%s", - provider, - model, - prompt_variant, - ) - raw_content = self._call_provider( - provider=provider, - model=model, - system_prompt=prompt["system"], - user_prompt=prompt["user"], - timeout=timeout, - ) - parsed = parse_verification_response(provider, raw_content) - if breaker: - breaker.record_success() - return { - "provider": provider, - "model": model, - "prompt_variant": prompt_variant, - "verification": parsed, - "raw_response": raw_content, - "stamp": { - "provider": provider, - "model": model, - "prompt_variant": prompt_variant, - } - } - except Exception as exc: - if breaker: - breaker.record_failure() - err = f"provider={provider}, model={model}, prompt={prompt_variant}, error={exc}" - errors.append(err) - logger.warning("Humanitarian verification attempt failed: %s", err) - - raise RuntimeError("All humanitarian verification attempts failed: " + " | ".join(errors)) - finally: - latency = time.time() - start_time - metrics.PIPELINE_STEP_LATENCY.labels(step_name='verify').observe(latency) - - def _provider_attempt_order(self, provider_preference: str) -> List[str]: - available: List[str] = [] - if settings.test_provider_mode: - available.append("test") - if settings.openai_api_key: - available.append("openai") - if settings.groq_api_key: - available.append("groq") - - preference = (provider_preference or "auto").lower() - if preference == "test" and settings.test_provider_mode: - return [preference] - if preference in ("openai", "groq", "test") and preference in available: - return [preference] + [provider for provider in available if provider != preference] - return available - - def _get_model_for_provider(self, provider: str) -> str: - if provider == "test": - return "test-provider/fixture" - if provider == "openai": - return settings.openai_model - if provider == "groq": - return settings.groq_model - raise ValueError(f"Unsupported provider: {provider}") - - def _call_provider( - self, - provider: str, - model: str, - system_prompt: str, - user_prompt: str, - timeout: Optional[float] = None, - ) -> str: - if provider == "test": - return self._call_test(model, system_prompt, user_prompt) - if provider == "openai": - return self._call_openai(model, system_prompt, user_prompt, timeout) - if provider == "groq": - return self._call_groq(model, system_prompt, user_prompt, timeout) - raise ValueError(f"Unsupported provider: {provider}") - - def _call_openai( - self, - model: str, - system_prompt: str, - user_prompt: str, - timeout: Optional[float] = None, - ) -> str: - if not settings.openai_api_key: - raise RuntimeError("OpenAI API key is not configured") - return self._call_chat_completion_api( - base_url="https://api.openai.com/v1/chat/completions", - api_key=settings.openai_api_key, - model=model, - system_prompt=system_prompt, - user_prompt=user_prompt, - timeout=timeout, - ) - - def _call_groq( - self, - model: str, - system_prompt: str, - user_prompt: str, - timeout: Optional[float] = None, - ) -> str: - if not settings.groq_api_key: - raise RuntimeError("Groq API key is not configured") - return self._call_chat_completion_api( - base_url="https://api.groq.com/openai/v1/chat/completions", - api_key=settings.groq_api_key, - model=model, - system_prompt=system_prompt, - user_prompt=user_prompt, - timeout=timeout, - ) - - def _call_chat_completion_api( - self, - base_url: str, - api_key: str, - model: str, - system_prompt: str, - user_prompt: str, - timeout: Optional[float] = None, - ) -> str: - if settings.ai_deterministic_mode: - logger.info("Deterministic AI mode enabled: returning stable response") - return self._get_deterministic_response(model, system_prompt, user_prompt) - - payload = { - "model": model, - "temperature": 0.1, - "messages": [ - {"role": "system", "content": system_prompt}, - {"role": "user", "content": user_prompt}, - ], - } - headers = { - "Authorization": f"Bearer {api_key}", - "Content-Type": "application/json", - } - - req_timeout = timeout if timeout is not None else float(settings.llm_timeout_seconds) - provider_name = "openai" if "openai" in base_url else "groq" - - try: - with httpx.Client(timeout=req_timeout) as client: - response = client.post(base_url, json=payload, headers=headers) - response.raise_for_status() - data = response.json() - except httpx.TimeoutException as exc: - logger.error("LLM provider %s request timed out after %s seconds", provider_name, req_timeout) - raise AIServiceError( - message=f"LLM request timed out after {req_timeout}s", - code="AI_TIMEOUT", - details={"provider": provider_name, "timeout_seconds": req_timeout}, - ) from exc - except httpx.HTTPStatusError as exc: - logger.error("LLM provider %s returned status %s: %s", provider_name, exc.response.status_code, exc.response.text) - raise AIServiceError( - message=f"LLM request failed with status {exc.response.status_code}", - code="AI_PROVIDER_ERROR", - details={"provider": provider_name, "status_code": exc.response.status_code}, - ) from exc - except Exception as exc: - logger.error("LLM provider %s connection or unexpected error: %s", provider_name, str(exc)) - raise AIServiceError( - message=f"LLM connection error: {str(exc)}", - code="AI_CONNECTION_ERROR", - details={"provider": provider_name}, - ) from exc - - try: - content = data["choices"][0]["message"]["content"] - except (KeyError, IndexError, TypeError) as exc: - raise RuntimeError(f"Unexpected LLM response format: {data}") from exc - - if not content: - raise RuntimeError("LLM returned empty content") - - return str(content) - - def _call_test(self, model: str, system_prompt: str, user_prompt: str) -> str: - response = self.test_provider.get_response( - endpoint="humanitarian", - request_data={ - "system_prompt": system_prompt, - "user_prompt": user_prompt, - }, - ) - return json.dumps(response, separators=(",", ":"), sort_keys=True) - - def _get_deterministic_response(self, model: str, system_prompt: str, user_prompt: str) -> str: - stable_response = { - "verdict": "credible", - "confidence": 0.74, - "summary": "Deterministic verification output for testing", - } - return json.dumps(stable_response, separators=(",", ":"), sort_keys=True) - - def _parse_json_response(self, content: str) -> Dict[str, Any]: - return parse_verification_response("auto", content) - - -def parse_verification_response(provider_name: str, raw_content: str) -> Dict[str, Any]: - """Parses raw verification response, handling JSON markdown blocks and potential truncations.""" - normalized = raw_content.strip() - if normalized.startswith("```"): - normalized = normalized.strip("`") - if normalized.startswith("json"): - normalized = normalized[4:].strip() - - try: - parsed = json.loads(normalized) - if isinstance(parsed, dict): - return _normalize_verification_dict(parsed) - except json.JSONDecodeError: - pass - - # Recovery parsing in case of truncation - import re - verdict_match = re.search(r'"verdict"\s*:\s*"([^"]+)"', normalized) - confidence_match = re.search(r'"confidence"\s*:\s*([0-9.]+)', normalized) - summary_match = re.search(r'"summary"\s*:\s*"([^"]*)"', normalized) - - verdict = verdict_match.group(1) if verdict_match else "inconclusive" - confidence = float(confidence_match.group(1)) if confidence_match else 0.0 - summary = summary_match.group(1) if summary_match else "Truncated response parsed via recovery" - - if verdict not in ["credible", "partially_credible", "inconclusive", "not_credible"]: - verdict = "inconclusive" - - return { - "verdict": verdict, - "confidence": confidence, - "summary": summary, - "criteria_assessment": None, - "risk_flags": None, - "missing_information": None, - "recommended_next_steps": None, - } - - -def _normalize_verification_dict(parsed: Dict[str, Any]) -> Dict[str, Any]: - """Ensures a parsed dict strictly matches HumanitarianVerificationDetailsV2 structure.""" - verdict = parsed.get("verdict", "inconclusive") - if verdict not in ["credible", "partially_credible", "inconclusive", "not_credible"]: - verdict = "inconclusive" - - confidence = parsed.get("confidence", 0.0) - try: - confidence = float(confidence) - except (ValueError, TypeError): - confidence = 0.0 - - return { - "verdict": verdict, - "confidence": confidence, - "summary": str(parsed.get("summary", "")), - "criteria_assessment": parsed.get("criteria_assessment"), - "risk_flags": parsed.get("risk_flags"), - "missing_information": parsed.get("missing_information"), - "recommended_next_steps": parsed.get("recommended_next_steps"), - } \ No newline at end of file +"""Humanitarian claim verification service with model/provider fallbacks.""" + +import json +import logging +from typing import Any, Dict, List, Optional +import time +import metrics + +import httpx + +from config import settings +from services.humanitarian_prompt import HumanitarianPromptEngine +from services.circuit_breaker import CircuitBreaker +from services.test_provider import TestProvider +from exceptions import AIServiceError + +logger = logging.getLogger(__name__) + + +class HumanitarianVerificationService: + """Runs humanitarian verification against configured LLM providers.""" + + def __init__(self): + self.prompt_engine = HumanitarianPromptEngine() + self.test_provider = TestProvider() + self.breakers = { + "openai": CircuitBreaker( + name="openai", + failure_threshold=settings.circuit_breaker_failure_threshold, + recovery_timeout=settings.circuit_breaker_recovery_timeout_seconds, + ), + "groq": CircuitBreaker( + name="groq", + failure_threshold=settings.circuit_breaker_failure_threshold, + recovery_timeout=settings.circuit_breaker_recovery_timeout_seconds, + ), + } + + def verify_claim( + self, + aid_claim: str, + supporting_evidence: Optional[List[str]] = None, + context_factors: Optional[Dict[str, Any]] = None, + provider_preference: str = "auto", + timeout: Optional[float] = None, + ) -> Dict[str, Any]: + start_time = time.time() + try: + evidence = supporting_evidence or [] + context = context_factors or {} + + primary_prompt = self.prompt_engine.build_primary_prompt( + aid_claim=aid_claim, + supporting_evidence=evidence, + context_factors=context, + ) + fallback_prompt = self.prompt_engine.build_fallback_prompt( + aid_claim=aid_claim, + supporting_evidence=evidence, + context_factors=context, + ) + + providers = self._provider_attempt_order(provider_preference) + if not providers: + raise RuntimeError("No LLM providers configured for humanitarian verification") + + errors: List[str] = [] + + for provider in providers: + breaker = self.breakers.get(provider) + if breaker and not breaker.allow_request(): + logger.warning("Circuit breaker is OPEN for provider=%s. Skipping.", provider) + errors.append(f"provider={provider}, error=Circuit breaker is OPEN") + continue + + model = self._get_model_for_provider(provider) + for prompt_variant, prompt in (("primary", primary_prompt), ("fallback", fallback_prompt)): + try: + logger.info( + "Attempting humanitarian verification with provider=%s model=%s prompt=%s", + provider, + model, + prompt_variant, + ) + raw_content = self._call_provider( + provider=provider, + model=model, + system_prompt=prompt["system"], + user_prompt=prompt["user"], + timeout=timeout, + ) + parsed = parse_verification_response(provider, raw_content) + if breaker: + breaker.record_success() + return { + "provider": provider, + "model": model, + "prompt_variant": prompt_variant, + "verification": parsed, + "raw_response": raw_content, + "stamp": { + "provider": provider, + "model": model, + "prompt_variant": prompt_variant, + } + } + except Exception as exc: + if breaker: + breaker.record_failure() + err = f"provider={provider}, model={model}, prompt={prompt_variant}, error={exc}" + errors.append(err) + logger.warning("Humanitarian verification attempt failed: %s", err) + + raise RuntimeError("All humanitarian verification attempts failed: " + " | ".join(errors)) + finally: + latency = time.time() - start_time + metrics.PIPELINE_STEP_LATENCY.labels(step_name='verify').observe(latency) + + def _provider_attempt_order(self, provider_preference: str) -> List[str]: + available: List[str] = [] + if settings.test_provider_mode: + available.append("test") + if settings.openai_api_key: + available.append("openai") + if settings.groq_api_key: + available.append("groq") + + preference = (provider_preference or "auto").lower() + if preference == "test" and settings.test_provider_mode: + return [preference] + if preference in ("openai", "groq", "test") and preference in available: + return [preference] + [provider for provider in available if provider != preference] + return available + + def _get_model_for_provider(self, provider: str) -> str: + if provider == "test": + return "test-provider/fixture" + if provider == "openai": + return settings.openai_model + if provider == "groq": + return settings.groq_model + raise ValueError(f"Unsupported provider: {provider}") + + def _call_provider( + self, + provider: str, + model: str, + system_prompt: str, + user_prompt: str, + timeout: Optional[float] = None, + ) -> str: + if provider == "test": + return self._call_test(model, system_prompt, user_prompt) + if provider == "openai": + return self._call_openai(model, system_prompt, user_prompt, timeout) + if provider == "groq": + return self._call_groq(model, system_prompt, user_prompt, timeout) + raise ValueError(f"Unsupported provider: {provider}") + + def _call_openai( + self, + model: str, + system_prompt: str, + user_prompt: str, + timeout: Optional[float] = None, + ) -> str: + if not settings.openai_api_key: + raise RuntimeError("OpenAI API key is not configured") + return self._call_chat_completion_api( + base_url="https://api.openai.com/v1/chat/completions", + api_key=settings.openai_api_key, + model=model, + system_prompt=system_prompt, + user_prompt=user_prompt, + timeout=timeout, + ) + + def _call_groq( + self, + model: str, + system_prompt: str, + user_prompt: str, + timeout: Optional[float] = None, + ) -> str: + if not settings.groq_api_key: + raise RuntimeError("Groq API key is not configured") + return self._call_chat_completion_api( + base_url="https://api.groq.com/openai/v1/chat/completions", + api_key=settings.groq_api_key, + model=model, + system_prompt=system_prompt, + user_prompt=user_prompt, + timeout=timeout, + ) + + def _call_chat_completion_api( + self, + base_url: str, + api_key: str, + model: str, + system_prompt: str, + user_prompt: str, + timeout: Optional[float] = None, + ) -> str: + if settings.ai_deterministic_mode: + logger.info("Deterministic AI mode enabled: returning stable response") + return self._get_deterministic_response(model, system_prompt, user_prompt) + + payload = { + "model": model, + "temperature": 0.1, + "messages": [ + {"role": "system", "content": system_prompt}, + {"role": "user", "content": user_prompt}, + ], + } + headers = { + "Authorization": f"Bearer {api_key}", + "Content-Type": "application/json", + } + + req_timeout = timeout if timeout is not None else float(settings.llm_timeout_seconds) + provider_name = "openai" if "openai" in base_url else "groq" + + try: + with httpx.Client(timeout=req_timeout) as client: + response = client.post(base_url, json=payload, headers=headers) + response.raise_for_status() + data = response.json() + except httpx.TimeoutException as exc: + logger.error("LLM provider %s request timed out after %s seconds", provider_name, req_timeout) + raise AIServiceError( + message=f"LLM request timed out after {req_timeout}s", + code="AI_TIMEOUT", + details={"provider": provider_name, "timeout_seconds": req_timeout}, + ) from exc + except httpx.HTTPStatusError as exc: + logger.error("LLM provider %s returned status %s: %s", provider_name, exc.response.status_code, exc.response.text) + raise AIServiceError( + message=f"LLM request failed with status {exc.response.status_code}", + code="AI_PROVIDER_ERROR", + details={"provider": provider_name, "status_code": exc.response.status_code}, + ) from exc + except Exception as exc: + logger.error("LLM provider %s connection or unexpected error: %s", provider_name, str(exc)) + raise AIServiceError( + message=f"LLM connection error: {str(exc)}", + code="AI_CONNECTION_ERROR", + details={"provider": provider_name}, + ) from exc + + try: + content = data["choices"][0]["message"]["content"] + except (KeyError, IndexError, TypeError) as exc: + raise RuntimeError(f"Unexpected LLM response format: {data}") from exc + + if not content: + raise RuntimeError("LLM returned empty content") + + return str(content) + + def _call_test(self, model: str, system_prompt: str, user_prompt: str) -> str: + response = self.test_provider.get_response( + endpoint="humanitarian", + request_data={ + "system_prompt": system_prompt, + "user_prompt": user_prompt, + }, + ) + return json.dumps(response, separators=(",", ":"), sort_keys=True) + + def _get_deterministic_response(self, model: str, system_prompt: str, user_prompt: str) -> str: + stable_response = { + "verdict": "credible", + "confidence": 0.74, + "summary": "Deterministic verification output for testing", + } + return json.dumps(stable_response, separators=(",", ":"), sort_keys=True) + + def _parse_json_response(self, content: str) -> Dict[str, Any]: + return parse_verification_response("auto", content) + + +def parse_verification_response(provider_name: str, raw_content: str) -> Dict[str, Any]: + """Parses raw verification response, handling JSON markdown blocks and potential truncations.""" + normalized = raw_content.strip() + if normalized.startswith("```"): + normalized = normalized.strip("`") + if normalized.startswith("json"): + normalized = normalized[4:].strip() + + try: + parsed = json.loads(normalized) + if isinstance(parsed, dict): + return _normalize_verification_dict(parsed) + except json.JSONDecodeError: + pass + + # Recovery parsing in case of truncation + import re + verdict_match = re.search(r'"verdict"\s*:\s*"([^"]+)"', normalized) + confidence_match = re.search(r'"confidence"\s*:\s*([0-9.]+)', normalized) + summary_match = re.search(r'"summary"\s*:\s*"([^"]*)"', normalized) + + verdict = verdict_match.group(1) if verdict_match else "inconclusive" + confidence = float(confidence_match.group(1)) if confidence_match else 0.0 + summary = summary_match.group(1) if summary_match else "Truncated response parsed via recovery" + + if verdict not in ["credible", "partially_credible", "inconclusive", "not_credible"]: + verdict = "inconclusive" + + return { + "verdict": verdict, + "confidence": confidence, + "summary": summary, + "criteria_assessment": None, + "risk_flags": None, + "missing_information": None, + "recommended_next_steps": None, + } + + +def _normalize_verification_dict(parsed: Dict[str, Any]) -> Dict[str, Any]: + """Ensures a parsed dict strictly matches HumanitarianVerificationDetailsV2 structure.""" + verdict = parsed.get("verdict", "inconclusive") + if verdict not in ["credible", "partially_credible", "inconclusive", "not_credible"]: + verdict = "inconclusive" + + confidence = parsed.get("confidence", 0.0) + try: + confidence = float(confidence) + except (ValueError, TypeError): + confidence = 0.0 + + return { + "verdict": verdict, + "confidence": confidence, + "summary": str(parsed.get("summary", "")), + "criteria_assessment": parsed.get("criteria_assessment"), + "risk_flags": parsed.get("risk_flags"), + "missing_information": parsed.get("missing_information"), + "recommended_next_steps": parsed.get("recommended_next_steps"), + } diff --git a/app/ai-service/tasks.py b/app/ai-service/tasks.py index 53f21fb4..fd4175ef 100644 --- a/app/ai-service/tasks.py +++ b/app/ai-service/tasks.py @@ -6,10 +6,12 @@ import logging import uuid import time +import json from typing import Any, Dict, Optional from celery import Celery from celery.result import AsyncResult import httpx +import redis import metrics from config import settings @@ -70,8 +72,32 @@ def process_heavy_inference_task(self, task_id: str, payload: Dict[str, Any]) -> return process_heavy_inference_task -# Task status storage (in production, use Redis with proper TTL) -task_results: Dict[str, Dict[str, Any]] = {} +# Lazy Redis client initialization - defers connection until needed +redis_client = None + +def get_redis_client() -> redis.Redis: + """ + Get or initialize the Redis client. + """ + global redis_client + if redis_client is None: + redis_client = redis.from_url(settings.redis_url, decode_responses=True) + return redis_client + + +def set_status(task_id: str, payload: Dict[str, Any]) -> None: + """ + Write task status payload directly to Redis with a 24-hour TTL. + """ + try: + r = get_redis_client() + key = f"task_status:{task_id}" + # TTL of 24 hours (86400 seconds) + r.setex(key, 86400, json.dumps(payload)) + except Exception as e: + logger.error(f"Failed to write task status to Redis: {e}") + + pii_scrubber_service = PIIScrubberService() humanitarian_verification_service = HumanitarianVerificationService() @@ -91,12 +117,13 @@ def update_task_status( result: Task result data (if completed) error: Error message (if failed) """ - task_results[task_id] = { + payload = { 'status': status, 'result': result, 'error': error, 'updated_at': time.time() } + set_status(task_id, payload) def send_webhook_notification(task_id: str, status: str, result: Any = None, error: str = None) -> None: @@ -373,20 +400,22 @@ def get_task_status(task_id: str) -> Dict[str, Any]: 'task_id': task_id, 'status': 'processing', } - else: - return { - 'task_id': task_id, - 'status': 'pending', - } except Exception: pass - # Fall back to local storage - if task_id in task_results: - return { - 'task_id': task_id, - **task_results[task_id] - } + # Fall back to Redis storage + try: + r = get_redis_client() + key = f"task_status:{task_id}" + data = r.get(key) + if data: + payload = json.loads(data) + return { + 'task_id': task_id, + **payload + } + except Exception as e: + logger.error(f"Failed to read task status from Redis: {e}") return { 'task_id': task_id, diff --git a/app/ai-service/tests/test_otel_spans.py b/app/ai-service/tests/test_otel_spans.py new file mode 100644 index 00000000..0eda1ed5 --- /dev/null +++ b/app/ai-service/tests/test_otel_spans.py @@ -0,0 +1,106 @@ +import os +import pytest +from unittest.mock import patch, MagicMock +from config import settings +from services.humanitarian_verification import HumanitarianVerificationService +from tracing.otel_setup import ( + reset_tracing_for_test, + get_in_memory_exporter, +) + +class TestOtelSpans: + @pytest.fixture(autouse=True) + def setup_otel(self, monkeypatch): + # Force app env to test to register InMemorySpanExporter + monkeypatch.setenv("APP_ENV", "test") + monkeypatch.setattr(settings, "openai_api_key", "test-openai-key") + monkeypatch.setattr(settings, "groq_api_key", "test-groq-key") + reset_tracing_for_test() + + self.exporter = get_in_memory_exporter() + if self.exporter: + self.exporter.clear() + + def test_verify_claim_emits_two_spans_openai(self, monkeypatch): + service = HumanitarianVerificationService() + + # Mock _call_chat_completion_api to avoid making real network requests + mock_response = '{"verdict": "credible", "confidence": 0.95, "summary": "verified"}' + monkeypatch.setattr( + service, + "_call_chat_completion_api", + lambda *args, **kwargs: mock_response + ) + + # Ensure we only try to call openai + monkeypatch.setattr(service, "_provider_attempt_order", lambda pref: ["openai"]) + monkeypatch.setattr(service, "_get_model_for_provider", lambda prov: "gpt-4-test") + + # Trigger claim verification + result = service.verify_claim( + aid_claim="Food packs delivered to flood zone.", + supporting_evidence=["waybill-102"], + context_factors={"weather": "clear"}, + provider_preference="openai" + ) + + assert result["provider"] == "openai" + assert result["prompt_variant"] == "primary" + + # Verify tracing spans + finished_spans = self.exporter.get_finished_spans() + assert len(finished_spans) == 2, f"Expected 2 spans, got {len(finished_spans)}" + + # Spans are emitted as they finish: + # call_openai is nested inside call_provider, so call_openai finishes first! + span_openai = finished_spans[0] + span_provider = finished_spans[1] + + assert span_openai.name == "humanitarian_verification.call_openai" + assert span_openai.attributes.get("model") == "gpt-4-test" + assert span_openai.attributes.get("prompt_variant") == "primary" + + assert span_provider.name == "humanitarian_verification.call_provider" + assert span_provider.attributes.get("model") == "gpt-4-test" + assert span_provider.attributes.get("prompt_variant") == "primary" + + def test_verify_claim_emits_two_spans_groq(self, monkeypatch): + service = HumanitarianVerificationService() + + # Mock _call_chat_completion_api to avoid making real network requests + mock_response = '{"verdict": "not_credible", "confidence": 0.85, "summary": "no evidence"}' + monkeypatch.setattr( + service, + "_call_chat_completion_api", + lambda *args, **kwargs: mock_response + ) + + # Ensure we only try to call groq + monkeypatch.setattr(service, "_provider_attempt_order", lambda pref: ["groq"]) + monkeypatch.setattr(service, "_get_model_for_provider", lambda prov: "llama3-groq-test") + + # Trigger claim verification + result = service.verify_claim( + aid_claim="Medicines delivered to shelter.", + supporting_evidence=["receipt-44"], + context_factors={"region": "north"}, + provider_preference="groq" + ) + + assert result["provider"] == "groq" + assert result["prompt_variant"] == "primary" + + # Verify tracing spans + finished_spans = self.exporter.get_finished_spans() + assert len(finished_spans) == 2, f"Expected 2 spans, got {len(finished_spans)}" + + span_groq = finished_spans[0] + span_provider = finished_spans[1] + + assert span_groq.name == "humanitarian_verification.call_groq" + assert span_groq.attributes.get("model") == "llama3-groq-test" + assert span_groq.attributes.get("prompt_variant") == "primary" + + assert span_provider.name == "humanitarian_verification.call_provider" + assert span_provider.attributes.get("model") == "llama3-groq-test" + assert span_provider.attributes.get("prompt_variant") == "primary" diff --git a/app/ai-service/tests/test_tasks_redis.py b/app/ai-service/tests/test_tasks_redis.py new file mode 100644 index 00000000..0690b70a --- /dev/null +++ b/app/ai-service/tests/test_tasks_redis.py @@ -0,0 +1,113 @@ +import json +import time +from unittest.mock import MagicMock, patch +import pytest +import tasks + +class MockRedis: + def __init__(self): + self.store = {} + self.ttls = {} + + def setex(self, key: str, time_to_live: int, value: str): + self.store[key] = value + self.ttls[key] = time.time() + time_to_live + + def get(self, key: str): + if key in self.store: + if time.time() < self.ttls[key]: + return self.store[key] + else: + del self.store[key] + del self.ttls[key] + return None + + def ttl(self, key: str): + if key in self.store: + remaining = self.ttls[key] - time.time() + return int(remaining) if remaining > 0 else -2 + return -2 + + +@pytest.fixture +def mock_redis(): + mr = MockRedis() + with patch("tasks.get_redis_client", return_value=mr): + # Reset lazy client to avoid caching previous states + with patch("tasks.redis_client", mr): + yield mr + + +def test_set_status_writes_to_redis_with_ttl(mock_redis): + task_id = "test-task-1" + payload = {"status": "processing", "result": None, "error": None} + + tasks.set_status(task_id, payload) + + key = f"task_status:{task_id}" + assert key in mock_redis.store + + stored_data = json.loads(mock_redis.store[key]) + assert stored_data["status"] == "processing" + + # Assert TTL is 24 hours (86400 seconds) + remaining_ttl = mock_redis.ttl(key) + assert 86300 <= remaining_ttl <= 86400 + + +def test_get_task_status_fallback_to_redis(mock_redis): + task_id = "test-task-2" + payload = {"status": "completed", "result": {"data": 123}, "error": None} + + # Populate redis + tasks.set_status(task_id, payload) + + # Celery raises Exception or returns non-ready task to trigger fallback + mock_async_result = MagicMock() + mock_async_result.ready.return_value = False + mock_async_result.started.return_value = False + + with patch("tasks.AsyncResult", return_value=mock_async_result): + status = tasks.get_task_status(task_id) + + assert status["task_id"] == task_id + assert status["status"] == "completed" + assert status["result"] == {"data": 123} + assert status["error"] is None + + +def test_get_task_status_celery_first(mock_redis): + task_id = "test-task-3" + + # Populate Redis with a different status + tasks.set_status(task_id, {"status": "processing", "result": None, "error": None}) + + # Celery returns ready task (completed) + mock_async_result = MagicMock() + mock_async_result.ready.return_value = True + mock_async_result.successful.return_value = True + mock_async_result.result = {"celery_data": 456} + + with patch("tasks.AsyncResult", return_value=mock_async_result): + status = tasks.get_task_status(task_id) + + # Should use Celery result, not Redis + assert status["status"] == "completed" + assert status["result"] == {"celery_data": 456} + + +def test_e2e_cross_process_observation(mock_redis): + # Simulate Process A updating the status + task_id = "cross-process-task-id" + tasks.update_task_status(task_id, "processing") + + # Simulate Process B retrieving the status + mock_async_result = MagicMock() + mock_async_result.ready.return_value = False + mock_async_result.started.return_value = False + + with patch("tasks.AsyncResult", return_value=mock_async_result): + status_b = tasks.get_task_status(task_id) + + assert status_b["task_id"] == task_id + assert status_b["status"] == "processing" diff --git a/app/ai-service/tracing/otel_setup.py b/app/ai-service/tracing/otel_setup.py new file mode 100644 index 00000000..cb8511b5 --- /dev/null +++ b/app/ai-service/tracing/otel_setup.py @@ -0,0 +1,64 @@ +import os +from opentelemetry import trace +from opentelemetry.sdk.trace import TracerProvider +from opentelemetry.sdk.trace.export import BatchSpanProcessor, SimpleSpanProcessor +from opentelemetry.sdk.trace.export.in_memory_span_exporter import InMemorySpanExporter +from opentelemetry.sdk.resources import Resource + +# Global variables to track state +_in_memory_exporter = None +_initialized = False + +def setup_tracing(): + global _in_memory_exporter, _initialized + if _initialized: + return + + resource = Resource.create(attributes={ + "service.name": "ai-service" + }) + + provider = TracerProvider(resource=resource) + trace.set_tracer_provider(provider) + + otel_endpoint = os.environ.get("OTEL_EXPORTER_OTLP_ENDPOINT") + app_env = os.environ.get("APP_ENV", "development") + + # In tests, prioritize InMemorySpanExporter so we can assert on spans + if app_env == "test": + _in_memory_exporter = InMemorySpanExporter() + # Use SimpleSpanProcessor for synchronous span processing in tests + provider.add_span_processor(SimpleSpanProcessor(_in_memory_exporter)) + elif otel_endpoint: + try: + from opentelemetry.exporter.otlp.proto.http.trace_exporter import OTLPSpanExporter + exporter = OTLPSpanExporter(endpoint=otel_endpoint) + # BatchSpanProcessor is standard for production OTLP exporting + provider.add_span_processor(BatchSpanProcessor(exporter)) + except Exception as e: + # Fallback to in-memory if OTLP setup fails + import logging + logging.getLogger(__name__).warning("Failed to initialize OTLPSpanExporter: %s. Falling back to InMemorySpanExporter.", e) + _in_memory_exporter = InMemorySpanExporter() + provider.add_span_processor(SimpleSpanProcessor(_in_memory_exporter)) + else: + # Default fallback (e.g. development without Jaeger) + _in_memory_exporter = InMemorySpanExporter() + provider.add_span_processor(SimpleSpanProcessor(_in_memory_exporter)) + + _initialized = True + +def get_tracer(): + if not _initialized: + setup_tracing() + return trace.get_tracer("ai-service") + +def get_in_memory_exporter(): + return _in_memory_exporter + +# Reset tracing state (mainly for clean unit testing) +def reset_tracing_for_test(): + global _in_memory_exporter, _initialized + _in_memory_exporter = None + _initialized = False + setup_tracing() diff --git a/app/backend/prisma/schema.prisma b/app/backend/prisma/schema.prisma index 454eb676..27c5f8e5 100644 --- a/app/backend/prisma/schema.prisma +++ b/app/backend/prisma/schema.prisma @@ -567,6 +567,7 @@ model RegistryOrganization { name String aliases String? // JSON array of alternative names externalId String? // External system identifier + provider String? metadata Json? createdAt DateTime @default(now()) updatedAt DateTime @updatedAt @@ -575,6 +576,7 @@ model RegistryOrganization { @@index([registryId]) @@index([name]) + @@unique([provider, externalId]) } /// Canonical registry for locations with stable IDs @@ -588,6 +590,7 @@ model RegistryLocation { coordinates Json? // { lat: number, lng: number } aliases String? // JSON array of alternative names externalId String? // External system identifier + provider String? metadata Json? createdAt DateTime @default(now()) updatedAt DateTime @updatedAt @@ -597,6 +600,7 @@ model RegistryLocation { @@index([registryId]) @@index([name]) @@index([country, region]) + @@unique([provider, externalId]) } /// Canonical registry for assets with stable IDs @@ -607,6 +611,7 @@ model RegistryAsset { type String? // e.g., "vehicle", "warehouse", "equipment" category String? externalId String? // External system identifier + provider String? metadata Json? createdAt DateTime @default(now()) updatedAt DateTime @updatedAt @@ -616,6 +621,7 @@ model RegistryAsset { @@index([registryId]) @@index([name]) @@index([type]) + @@unique([provider, externalId]) } /// Canonical registry for projects with stable IDs @@ -628,6 +634,7 @@ model RegistryProject { startDate DateTime? endDate DateTime? externalId String? // External system identifier + provider String? metadata Json? createdAt DateTime @default(now()) updatedAt DateTime @updatedAt @@ -637,6 +644,7 @@ model RegistryProject { @@index([registryId]) @@index([name]) @@index([status]) + @@unique([provider, externalId]) } enum EntityLinkSourceType {