diff --git a/backend/api/extractions/router.py b/backend/api/extractions/router.py index 1252f73..e4dd270 100644 --- a/backend/api/extractions/router.py +++ b/backend/api/extractions/router.py @@ -39,6 +39,7 @@ "azure-llama": "azure", "macbook": "macbook", "vllm": "vllm", + "cohere": "cohere", } # Timeout logging setup diff --git a/backend/api/server/router.py b/backend/api/server/router.py index 95d81e1..3183c97 100644 --- a/backend/api/server/router.py +++ b/backend/api/server/router.py @@ -577,6 +577,21 @@ async def get_available_models(): except Exception as e: print(f"[VLLM] Failed to fetch models: {e}") + if os.getenv("COHERE_AZURE_ENDPOINT") and os.getenv("COHERE_AZURE_KEY"): + models.append( + { + "id": os.getenv("COHERE_MODEL_NAME", "cohere-command-a"), + "name": os.getenv("COHERE_DISPLAY_NAME", "Cohere Command A"), + "provider": "Cohere", + "model_type": "cohere", + "description": "Azure AI Foundry", + "supports_temperature": True, + "default_temperature": 0.5, + "vision_capable": False, + } + ) + print("✅ Cohere Command A added to model list") + return JSONResponse(status_code=200, content=models) diff --git a/backend/core/config.py b/backend/core/config.py index ba32cee..2daa815 100644 --- a/backend/core/config.py +++ b/backend/core/config.py @@ -130,6 +130,30 @@ def load_config(): if anthropic_location: os.environ.setdefault("ANTHROPIC_LOCATION", anthropic_location) + # Cohere configuration (Azure AI Foundry) + cohere_cfg = cfg.get("cohere", {}) or {} + if cohere_cfg.get("endpoint"): + os.environ.setdefault( + "COHERE_AZURE_ENDPOINT", cohere_cfg["endpoint"] + ) + print( + f"✅ Cohere endpoint loaded from secrets.toml: {cohere_cfg['endpoint']}" + ) + if cohere_cfg.get("api_key"): + os.environ.setdefault("COHERE_AZURE_KEY", cohere_cfg["api_key"]) + os.environ.setdefault( + "COHERE_AZURE_API_VERSION", + cohere_cfg.get("api_version", "2024-05-01-preview"), + ) + os.environ.setdefault( + "COHERE_MODEL_NAME", + cohere_cfg.get("model_name", "cohere-command-a"), + ) + os.environ.setdefault( + "COHERE_DISPLAY_NAME", + cohere_cfg.get("display_name", "Cohere Command A"), + ) + # Set up Google Cloud credentials for Vertex AI (shared by Gemini and Anthropic) # Look for service account key in backend/core/ directory service_account_path = ( diff --git a/backend/services/chat_memory/chat_memory_service.py b/backend/services/chat_memory/chat_memory_service.py index 37940ce..db65685 100644 --- a/backend/services/chat_memory/chat_memory_service.py +++ b/backend/services/chat_memory/chat_memory_service.py @@ -739,6 +739,8 @@ def _estimate_max_context_tokens( return 128_000 if model_type in ("macbook", "vllm"): return 32_000 + if model_type == "cohere": + return 256_000 # Command A context window if "gpt54mini" in compact_model_key or "gpt54nano" in compact_model_key: return 400_000 if "gpt54" in compact_model_key: diff --git a/backend/services/llm/cohere.py b/backend/services/llm/cohere.py new file mode 100644 index 0000000..57131dd --- /dev/null +++ b/backend/services/llm/cohere.py @@ -0,0 +1,296 @@ +""" +Cohere LLM Client — calls Cohere Command A via Azure AI Foundry. + +Supports two Azure AI endpoint types (auto-detected from the URL): + - Azure AI Services hub (*.services.ai.azure.com): + POST {base}/models/chat/completions?api-version={version} + Header: api-key: {key} + Body: "model": {model_name} + - Azure AI Serverless (*.models.ai.azure.com): + POST {base}/v1/chat/completions + Header: Authorization: Bearer {key} + +Configuration (from secrets.toml [cohere] section, loaded into env vars): + COHERE_AZURE_ENDPOINT — base URL of the Azure AI Foundry endpoint + COHERE_AZURE_KEY — API key + COHERE_AZURE_API_VERSION — e.g. "2024-05-01-preview" (hub only) + COHERE_MODEL_NAME — deployment/model name, e.g. "cohere-command-a" +""" + +import asyncio +import json +import os +import time +from typing import Any, Dict, List, Optional +from urllib.parse import urlparse + +import requests + +_DEFAULT_EXTRACTION_SYSTEM = ( + "You are an expert scientific data extractor. " + "Respond ONLY with a valid JSON object in this exact format:\n" + '{"answer": "", "references": [{"text": ""}]}\n' + "Include one or more reference objects quoting the specific passages you used as evidence. " + "No markdown fences, no extra keys." +) + + +class CohereLLMClient: + def __init__(self): + raw_endpoint = os.environ.get("COHERE_AZURE_ENDPOINT", "") + self.api_key = os.environ.get("COHERE_AZURE_KEY", "") + self.api_version = os.environ.get( + "COHERE_AZURE_API_VERSION", "2024-05-01-preview" + ) + self.model_name = os.environ.get("COHERE_MODEL_NAME", "cohere-command-a") + + # Strip path components — always work from the bare host + parsed = urlparse(raw_endpoint) + self.endpoint = ( + f"{parsed.scheme}://{parsed.netloc}" + if parsed.netloc + else raw_endpoint.rstrip("/") + ) + + self.disabled = not bool(self.endpoint and self.api_key) + + _host = urlparse(self.endpoint).hostname or "" + self._is_hub = _host.endswith(".services.ai.azure.com") + self._is_serverless = _host.endswith(".models.ai.azure.com") + + if not self.disabled: + kind = ( + "hub" + if self._is_hub + else ("serverless" if self._is_serverless else "unknown") + ) + print(f"[CohereLLMClient] Initialised → {self.endpoint} ({kind})") + else: + print( + "[CohereLLMClient] Disabled (COHERE_AZURE_ENDPOINT or COHERE_AZURE_KEY not set)" + ) + + def _build_url(self) -> str: + if self._is_hub: + return f"{self.endpoint}/models/chat/completions?api-version={self.api_version}" + return f"{self.endpoint}/v1/chat/completions" + + def _build_headers(self) -> Dict[str, str]: + headers = {"Content-Type": "application/json"} + if self._is_hub: + headers["api-key"] = self.api_key + else: + headers["Authorization"] = f"Bearer {self.api_key}" + return headers + + async def _call_api( + self, + messages: List[Dict[str, str]], + max_tokens: int = 8048, + temperature: Optional[float] = 0.5, + ) -> Dict[str, Any]: + """Core HTTP call with retry logic (3 attempts, exponential backoff).""" + url = self._build_url() + headers = self._build_headers() + + payload: Dict[str, Any] = { + "messages": messages, + "max_tokens": max_tokens, + } + if temperature is not None: + payload["temperature"] = temperature + if self._is_hub: + payload["model"] = self.model_name + + max_retries = 3 + base_delay = 1.0 + max_delay = 30.0 + retryable = {429, 500, 503, 504} + last_error: Optional[str] = None + + for attempt in range(max_retries): + try: + t0 = time.perf_counter() + resp = await asyncio.to_thread( + lambda: requests.post( + url, json=payload, headers=headers, timeout=120 + ) + ) + duration = time.perf_counter() - t0 + + if resp.status_code == 200: + data = resp.json() + choices = data.get("choices", []) + if not choices: + return { + "success": False, + "error": "No choices in Cohere response", + } + content = choices[0].get("message", {}).get("content", "") + usage = data.get("usage", {}) + return { + "success": True, + "content": content, + "meta": { + "model": self.model_name, + "provider": "cohere", + "prompt_tokens": usage.get("prompt_tokens", 0), + "completion_tokens": usage.get("completion_tokens", 0), + "total_tokens": usage.get("total_tokens", 0), + "duration": duration, + }, + } + + if resp.status_code in retryable and attempt < max_retries - 1: + import random + + delay = min( + base_delay * (2**attempt) + random.uniform(0, 1), max_delay + ) + print( + f"[CohereLLMClient] HTTP {resp.status_code} on attempt {attempt + 1}, retrying in {delay:.1f}s" + ) + await asyncio.sleep(delay) + continue + + # Non-retryable or last attempt + try: + err_body = resp.json() + except Exception: + err_body = resp.text[:500] + last_error = f"Azure API error (status {resp.status_code}): {err_body}" + if resp.status_code == 404: + last_error += ( + f"\n⚠️ Model '{self.model_name}' not found at {self.endpoint}. " + "Verify the deployment name and endpoint in secrets.toml." + ) + break + + except requests.exceptions.Timeout: + last_error = f"Cohere request timed out (attempt {attempt + 1})" + if attempt < max_retries - 1: + await asyncio.sleep(base_delay * (2**attempt)) + continue + break + except Exception as exc: + last_error = f"Cohere request error: {str(exc)}" + break + + return {"success": False, "error": last_error or "Unknown error"} + + def _parse_json_content(self, content: str) -> Dict[str, Any]: + """Parse JSON from model response, extracting answer and references fields.""" + stripped = content.strip() + if stripped.startswith("```"): + lines = stripped.split("\n") + stripped = "\n".join(lines[1:-1]).strip() if len(lines) > 2 else stripped + + try: + parsed = json.loads(stripped) + if isinstance(parsed, dict): + answer = parsed.get("answer") or stripped + references = parsed.get("references", []) + if isinstance(references, list): + references = [ + r if isinstance(r, dict) else {"text": str(r)} + for r in references + ] + return { + "success": True, + "extracted_text": answer, + "references": references, + } + # json.loads returned a scalar or list — treat as plain text + except json.JSONDecodeError: + # Non-JSON response is expected when the model returns plain text; fall through to raw return below + pass + + return {"success": True, "extracted_text": content, "references": []} + + async def extract_entities_with_cohere( + self, + markdown: str, + extraction_prompt: str, + max_tokens: int = 8048, + temperature: float = 0.0, + system_message: Optional[str] = None, + ) -> Dict[str, Any]: + if self.disabled: + return {"success": False, "error": "Cohere is not configured."} + + effective_system = system_message or _DEFAULT_EXTRACTION_SYSTEM + user_content = ( + f"{extraction_prompt}\n\n" + "IMPORTANT: Respond ONLY with a valid JSON object: " + '{"answer": "...", "references": [{"text": "exact verbatim quote"}]}\n\n' + f"---\n\n{markdown}" + ) + messages: List[Dict[str, str]] = [ + {"role": "system", "content": effective_system}, + {"role": "user", "content": user_content}, + ] + + result = await self._call_api( + messages, max_tokens=max_tokens, temperature=temperature + ) + if not result["success"]: + return result + + meta = result["meta"] + parsed = self._parse_json_content(result["content"]) + + return { + "success": True, + "content": parsed["extracted_text"], + "references": parsed["references"], + "model": self.model_name, + "meta": { + "model_name": self.model_name, + "deployment": self.model_name, + "prompt_tokens": meta.get("prompt_tokens", 0), + "completion_tokens": meta.get("completion_tokens", 0), + "total_tokens": meta.get("total_tokens", 0), + "duration": meta.get("duration", 0), + }, + } + + async def generate_paragraph_with_cohere( + self, + user_prompt: str, + max_tokens: int = 2048, + temperature: Optional[float] = 0.5, + system_message: Optional[str] = None, + ) -> Dict[str, Any]: + if self.disabled: + return {"success": False, "error": "Cohere is not configured."} + + default_system = ( + "You are a scientific writing assistant. Your task is to synthesize extracted " + "information into a cohesive, well-structured paragraph while maintaining complete " + "accuracy. Follow the instructions exactly and preserve all factual details." + ) + messages: List[Dict[str, str]] = [ + {"role": "system", "content": system_message or default_system}, + {"role": "user", "content": user_prompt}, + ] + + result = await self._call_api( + messages, max_tokens=max_tokens, temperature=temperature + ) + if not result["success"]: + return result + + meta = result["meta"] + return { + "success": True, + "content": result["content"], + "model": self.model_name, + "meta": { + "model_name": self.model_name, + "deployment": self.model_name, + "prompt_tokens": meta.get("prompt_tokens", 0), + "completion_tokens": meta.get("completion_tokens", 0), + "total_tokens": meta.get("total_tokens", 0), + "duration": meta.get("duration", 0), + }, + } diff --git a/backend/services/llm/llm_service.py b/backend/services/llm/llm_service.py index 64bdb4e..79e5e93 100644 --- a/backend/services/llm/llm_service.py +++ b/backend/services/llm/llm_service.py @@ -4,6 +4,7 @@ from pathlib import Path from typing import Dict, Any, Optional from .azure import AzureLLMClient +from .cohere import CohereLLMClient from .gemini import GeminiLLMClient from .anthropic import AnthropicLLMClient from .llama import LlamaLLMClient @@ -37,6 +38,7 @@ class LLMService: def __init__(self): self.azure_client = AzureLLMClient() + self.cohere_client = CohereLLMClient() self.gemini_client = GeminiLLMClient() self.anthropic_client = AnthropicLLMClient() self.llama_client = LlamaLLMClient() @@ -284,6 +286,21 @@ async def extract_entities_from_markdown( ) self._record_session_metrics(session_id, "vllm", result) return result + elif model_type == "cohere": + if self.cohere_client.disabled: + return {"success": False, "error": "Cohere is not configured."} + result = await self._call_with_timeout_logging( + operation_name, + self.cohere_client.extract_entities_with_cohere( + markdown, + extraction_prompt, + max_tokens, + temperature, + system_message, + ), + ) + self._record_session_metrics(session_id, "cohere", result) + return result else: return { "success": False, @@ -473,6 +490,20 @@ async def generate_paragraph( ) self._record_session_metrics(session_id, "vllm", result) return result + elif model_type == "cohere": + if self.cohere_client.disabled: + return {"success": False, "error": "Cohere is not configured."} + result = await self._call_with_timeout_logging( + "paragraph_cohere", + self.cohere_client.generate_paragraph_with_cohere( + user_prompt, + max_tokens, + temperature, + system_message, + ), + ) + self._record_session_metrics(session_id, "cohere", result) + return result else: return {"success": False, "error": f"Unsupported model type: {model_type}"} diff --git a/backend/tests/test_chat_memory_service.py b/backend/tests/test_chat_memory_service.py index cf26a91..cc9aeee 100644 --- a/backend/tests/test_chat_memory_service.py +++ b/backend/tests/test_chat_memory_service.py @@ -186,3 +186,20 @@ def test_context_usage_reports_small_nonzero_percentages(): assert context_usage["estimated_tokens"] > 0 assert 0 < context_usage["percentage"] < 0.1 + + +def test_context_window_cohere(): + service = ChatMemoryService(use_memory_checkpointer=True) + + context_usage = service._build_context_usage( + user_prompt="test", + system_message="", + messages=[], + conversation_summary="", + document_context=None, + model_type="cohere", + model_id=None, + deployment=None, + ) + + assert context_usage["max_tokens"] == 256_000 diff --git a/docker-compose.yml b/docker-compose.yml index f0ce2ba..29a7864 100644 --- a/docker-compose.yml +++ b/docker-compose.yml @@ -18,7 +18,7 @@ services: azurite: image: mcr.microsoft.com/azure-storage/azurite # Blob-only mode; --loose relaxes strict account validation for dev - command: azurite-blob --blobHost 0.0.0.0 --loose + command: azurite-blob --blobHost 0.0.0.0 --loose --skipApiVersionCheck ports: - "10000:10000" volumes: @@ -54,8 +54,22 @@ services: WORKERS: "1" # Azurite local dev account — this is a publicly documented development key, not a real secret. # See: https://learn.microsoft.com/azure/storage/common/storage-use-azurite#well-known-storage-account-and-key - AZURE_STORAGE_CONNECTION_STRING: "DefaultEndpointsProtocol=http;AccountName=devstoreaccount1;AccountKey=;BlobEndpoint=http://azurite:10000/devstoreaccount1;" + AZURE_STORAGE_CONNECTION_STRING: "DefaultEndpointsProtocol=http;AccountName=devstoreaccount1;AccountKey=;BlobEndpoint=http://azurite:10000/devstoreaccount1;" AZURE_STORAGE_CONTAINER_NAME: summarization-uploads + # Azure OpenAI — sourced from .env + AZURE_OPENAI_ENDPOINT: ${AZURE_OPENAI_ENDPOINT:-} + AZURE_OPENAI_KEY: ${AZURE_OPENAI_KEY:-} + AZURE_OPENAI_DEPLOYMENT: ${AZURE_OPENAI_DEPLOYMENT:-} + AZURE_OPENAI_MODEL_NAME: ${AZURE_OPENAI_MODEL_NAME:-} + AZURE_OPENAI_API_VERSION: ${AZURE_OPENAI_API_VERSION:-2025-04-01-preview} + # Cohere — sourced from .env + COHERE_AZURE_ENDPOINT: ${COHERE_AZURE_ENDPOINT:-} + COHERE_AZURE_KEY: ${COHERE_AZURE_KEY:-} + COHERE_AZURE_API_VERSION: ${COHERE_AZURE_API_VERSION:-2024-05-01-preview} + COHERE_MODEL_NAME: ${COHERE_MODEL_NAME:-cohere-command-a} + # Azure Document Intelligence — sourced from .env + AZURE_DOC_INTELLIGENCE_ENDPOINT: ${AZURE_DOC_INTELLIGENCE_ENDPOINT:-} + AZURE_DOC_INTELLIGENCE_KEY: ${AZURE_DOC_INTELLIGENCE_KEY:-} ports: - "8001:8001" depends_on: diff --git a/frontend/components/BatchStudySelectionPage.tsx b/frontend/components/BatchStudySelectionPage.tsx index dc1e306..e3f620f 100644 --- a/frontend/components/BatchStudySelectionPage.tsx +++ b/frontend/components/BatchStudySelectionPage.tsx @@ -834,6 +834,88 @@ export function BatchStudySelectionPage({ ); })()} + {/* Cohere Models */} + {(() => { + const cohereModels = availableModels.filter((m) => + m.provider?.toLowerCase().includes("cohere") + ); + if (cohereModels.length === 0) return null; + + const cohereIds = cohereModels.map((m) => m.id); + const selection = getCategorySelection(cohereIds); + + return ( + + +
+
+ Cohere + + {selection.selected}/{selection.total} + +
+
+
+ +
+
+ + {cohereModels.length} model + {cohereModels.length !== 1 ? "s" : ""} available + + +
+
+ {cohereModels.map((model) => ( +
+ toggleModel(model.id)} + className="mt-1" + /> +
+ +

+ {model.description || model.provider} +

+
+
+ ))} +
+
+
+
+ ); + })()} + {/* Macbook LLM Models */} {(() => { const macbookModels = availableModels.filter((m) => diff --git a/frontend/components/ChatPage.tsx b/frontend/components/ChatPage.tsx index d5f95df..70c857b 100644 --- a/frontend/components/ChatPage.tsx +++ b/frontend/components/ChatPage.tsx @@ -602,6 +602,7 @@ export function ChatPage({ onSwitchToWorkflow, onSignOut }: ChatPageProps) { if (!model) return null; const isGemini = model.provider === "Google Gemini"; const isAnthropic = model.provider === "Anthropic"; + const isCohere = model.provider === "Cohere"; const chatModel = model as ChatModelConfig; const isLlama = model.provider === "Meta Llama" || chatModel.model_type === "azure-llama"; @@ -612,7 +613,9 @@ export function ChatPage({ onSwitchToWorkflow, onSignOut }: ChatPageProps) { ? "anthropic" : isLlama ? "llama" - : "azure", + : isCohere + ? "cohere" + : "azure", modelId: model.id, deployment: model.deployment, apiVersion: model.api_version, diff --git a/frontend/components/EntityExtractionPage.tsx b/frontend/components/EntityExtractionPage.tsx index 732d9e9..ee5bb92 100644 --- a/frontend/components/EntityExtractionPage.tsx +++ b/frontend/components/EntityExtractionPage.tsx @@ -1792,6 +1792,8 @@ export function EntityExtractionPage({ modelType = "gemini"; } else if (provider.includes("anthropic")) { modelType = "anthropic"; + } else if (provider.includes("cohere")) { + modelType = "cohere"; } else if (provider.includes("azure")) { modelType = "azure"; } else if (provider.includes("meta") || provider.includes("llama")) { @@ -1871,7 +1873,7 @@ export function EntityExtractionPage({ return { modelId, result: null }; } - // Map provider to backend model_type ("azure", "gemini", "anthropic", "llama") + // Map provider to backend model_type ("azure", "gemini", "anthropic", "cohere", "llama") let modelType = "azure"; const provider = modelObj.provider?.toLowerCase() || ""; @@ -1879,6 +1881,8 @@ export function EntityExtractionPage({ modelType = "gemini"; } else if (provider.includes("anthropic")) { modelType = "anthropic"; + } else if (provider.includes("cohere")) { + modelType = "cohere"; } else if (provider.includes("azure")) { modelType = "azure"; } else if (provider.includes("meta") || provider.includes("llama")) { @@ -2380,6 +2384,7 @@ export function EntityExtractionPage({ const p = providerStr.toLowerCase(); if (p.includes("google") || p.includes("gemini")) return "gemini"; if (p.includes("anthropic")) return "anthropic"; + if (p.includes("cohere")) return "cohere"; if (p.includes("meta") || p.includes("llama")) return modelId?.startsWith("azure-") ? "azure-llama" : "llama"; if (p.includes("macbook")) return "macbook"; diff --git a/frontend/utils/modelSelection.ts b/frontend/utils/modelSelection.ts index 3a60dfd..747b824 100644 --- a/frontend/utils/modelSelection.ts +++ b/frontend/utils/modelSelection.ts @@ -79,6 +79,11 @@ const MODEL_PRIORITY: Array<{ match: (m) => m.id?.includes("llama-4-maverick"), modelType: "llama", }, + // Tier 6 — Cohere + { + match: (m) => m.provider === "Cohere", + modelType: "cohere", + }, // Catch-all { match: (m) => m.provider === "Google Gemini", @@ -157,13 +162,16 @@ export function modelConfigToSelection( const isLlama = model.provider === "Meta Llama" || (model as any).model_type === "azure-llama"; + const isCohere = model.provider === "Cohere"; const modelType = isGemini ? "gemini" : isAnthropic ? "anthropic" : isLlama ? "azure-llama" - : "azure"; + : isCohere + ? "cohere" + : "azure"; return { model, modelType,