From 1454755f3567ebb4718c60d40149811b567df415 Mon Sep 17 00:00:00 2001 From: Jordan Leis Date: Fri, 19 Jun 2026 13:57:14 +0000 Subject: [PATCH 1/5] feat: enable Cohere Command A in chat and advanced mode MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Ports the Cohere LLM client from the cohere-router branch and wires it into main. The model was already provisioned in Azure AI Foundry but never appeared in the app because the config loading, client, and model list endpoint were never merged. - Add CohereLLMClient (Azure AI Hub + Serverless endpoint auto-detection, 3-attempt exponential backoff, asyncio.to_thread for blocking HTTP) - Load [cohere] section from secrets.toml into env vars on startup - Route model_type="cohere" in both extract_entities and generate_paragraph - Expose Cohere in /api/models when COHERE_AZURE_ENDPOINT + KEY are set - Add "cohere" → 256K to chat memory context window estimator (Command A) - Map provider="Cohere" → model_type="cohere" in ChatPage getModelConfig - Add Cohere as Tier 6 in auto-selection priority (modelSelection.ts) - Add "cohere" to extraction provider cost tracking map --- backend/api/extractions/router.py | 1 + backend/api/server/router.py | 15 + backend/core/config.py | 11 + .../chat_memory/chat_memory_service.py | 2 + backend/services/llm/cohere.py | 263 ++++++++++++++++++ backend/services/llm/llm_service.py | 31 +++ frontend/components/ChatPage.tsx | 5 +- frontend/utils/modelSelection.ts | 5 + 8 files changed, 332 insertions(+), 1 deletion(-) create mode 100644 backend/services/llm/cohere.py 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..108e247 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": "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..a9c20c1 100644 --- a/backend/core/config.py +++ b/backend/core/config.py @@ -130,6 +130,17 @@ 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..180c015 --- /dev/null +++ b/backend/services/llm/cohere.py @@ -0,0 +1,263 @@ +""" +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) + + self._is_hub = ".services.ai.azure.com" in self.endpoint + self._is_serverless = ".models.ai.azure.com" in self.endpoint + + 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") + time.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: + time.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: + 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/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/utils/modelSelection.ts b/frontend/utils/modelSelection.ts index 3a60dfd..e32cf5a 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", From e14cda3064b97522873204db123119e5c3b9a4ba Mon Sep 17 00:00:00 2001 From: Jordan Leis Date: Fri, 19 Jun 2026 17:58:00 +0000 Subject: [PATCH 2/5] fix: wire Cohere Command A into entity extraction UI and fix local dev stack MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit - EntityExtractionPage: add "cohere" branch to all three provider→model_type mappings so Cohere no longer silently routes as Azure during extraction - BatchStudySelectionPage: add missing Cohere accordion section to the AI Models Selection panel - docker-compose.yml: fix placeholder with real well-known Azurite dev key; add --skipApiVersionCheck; wire AZURE_OPENAI_*, COHERE_*, and AZURE_DOC_INTELLIGENCE_* env vars to backend service via .env substitution --- docker-compose.yml | 18 +++- .../components/BatchStudySelectionPage.tsx | 82 +++++++++++++++++++ frontend/components/EntityExtractionPage.tsx | 7 +- 3 files changed, 104 insertions(+), 3 deletions(-) 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/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"; From 286322e9e08d1d9eb338c20d40e302a32080120e Mon Sep 17 00:00:00 2001 From: "github-actions[bot]" Date: Fri, 19 Jun 2026 13:59:07 +0000 Subject: [PATCH 3/5] style: auto-format with Black --- backend/core/config.py | 23 ++++++++++--- backend/services/llm/cohere.py | 59 ++++++++++++++++++++++++++-------- 2 files changed, 63 insertions(+), 19 deletions(-) diff --git a/backend/core/config.py b/backend/core/config.py index a9c20c1..2daa815 100644 --- a/backend/core/config.py +++ b/backend/core/config.py @@ -133,13 +133,26 @@ def load_config(): # 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']}") + 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")) + 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 diff --git a/backend/services/llm/cohere.py b/backend/services/llm/cohere.py index 180c015..8a008f8 100644 --- a/backend/services/llm/cohere.py +++ b/backend/services/llm/cohere.py @@ -39,12 +39,18 @@ 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.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.endpoint = ( + f"{parsed.scheme}://{parsed.netloc}" + if parsed.netloc + else raw_endpoint.rstrip("/") + ) self.disabled = not bool(self.endpoint and self.api_key) @@ -52,10 +58,16 @@ def __init__(self): self._is_serverless = ".models.ai.azure.com" in self.endpoint if not self.disabled: - kind = "hub" if self._is_hub else ("serverless" if self._is_serverless else "unknown") + 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)") + print( + "[CohereLLMClient] Disabled (COHERE_AZURE_ENDPOINT or COHERE_AZURE_KEY not set)" + ) def _build_url(self) -> str: if self._is_hub: @@ -99,7 +111,9 @@ async def _call_api( try: t0 = time.perf_counter() resp = await asyncio.to_thread( - lambda: requests.post(url, json=payload, headers=headers, timeout=120) + lambda: requests.post( + url, json=payload, headers=headers, timeout=120 + ) ) duration = time.perf_counter() - t0 @@ -107,7 +121,10 @@ async def _call_api( data = resp.json() choices = data.get("choices", []) if not choices: - return {"success": False, "error": "No choices in Cohere response"} + return { + "success": False, + "error": "No choices in Cohere response", + } content = choices[0].get("message", {}).get("content", "") usage = data.get("usage", {}) return { @@ -125,8 +142,13 @@ async def _call_api( 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") + + 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" + ) time.sleep(delay) continue @@ -146,7 +168,7 @@ async def _call_api( except requests.exceptions.Timeout: last_error = f"Cohere request timed out (attempt {attempt + 1})" if attempt < max_retries - 1: - time.sleep(base_delay * (2 ** attempt)) + time.sleep(base_delay * (2**attempt)) continue break except Exception as exc: @@ -169,9 +191,14 @@ def _parse_json_content(self, content: str) -> Dict[str, Any]: references = parsed.get("references", []) if isinstance(references, list): references = [ - r if isinstance(r, dict) else {"text": str(r)} for r in references + r if isinstance(r, dict) else {"text": str(r)} + for r in references ] - return {"success": True, "extracted_text": answer, "references": references} + return { + "success": True, + "extracted_text": answer, + "references": references, + } # json.loads returned a scalar or list — treat as plain text except json.JSONDecodeError: pass @@ -192,7 +219,7 @@ async def extract_entities_with_cohere( effective_system = system_message or _DEFAULT_EXTRACTION_SYSTEM user_content = ( f"{extraction_prompt}\n\n" - 'IMPORTANT: Respond ONLY with a valid JSON object: ' + "IMPORTANT: Respond ONLY with a valid JSON object: " '{"answer": "...", "references": [{"text": "exact verbatim quote"}]}\n\n' f"---\n\n{markdown}" ) @@ -201,7 +228,9 @@ async def extract_entities_with_cohere( {"role": "user", "content": user_content}, ] - result = await self._call_api(messages, max_tokens=max_tokens, temperature=temperature) + result = await self._call_api( + messages, max_tokens=max_tokens, temperature=temperature + ) if not result["success"]: return result @@ -243,7 +272,9 @@ async def generate_paragraph_with_cohere( {"role": "user", "content": user_prompt}, ] - result = await self._call_api(messages, max_tokens=max_tokens, temperature=temperature) + result = await self._call_api( + messages, max_tokens=max_tokens, temperature=temperature + ) if not result["success"]: return result From 70a4f5fc0401546964cd2b7ae4fa0b10091d85b7 Mon Sep 17 00:00:00 2001 From: Jordan Leis Date: Fri, 19 Jun 2026 18:16:43 +0000 Subject: [PATCH 4/5] fix: address CodeQL and code-quality findings in CohereLLMClient - Use urlparse().hostname + endswith() for hub/serverless detection instead of substring `in` check, preventing spoofed URLs from matching (CodeQL CWE-020 / incomplete URL sanitization) - Add explanatory comment to bare except json.JSONDecodeError pass so intent is clear (plain-text fallback is deliberate) --- backend/services/llm/cohere.py | 6 ++++-- 1 file changed, 4 insertions(+), 2 deletions(-) diff --git a/backend/services/llm/cohere.py b/backend/services/llm/cohere.py index 8a008f8..35a67d2 100644 --- a/backend/services/llm/cohere.py +++ b/backend/services/llm/cohere.py @@ -54,8 +54,9 @@ def __init__(self): self.disabled = not bool(self.endpoint and self.api_key) - self._is_hub = ".services.ai.azure.com" in self.endpoint - self._is_serverless = ".models.ai.azure.com" in self.endpoint + _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 = ( @@ -201,6 +202,7 @@ def _parse_json_content(self, content: str) -> Dict[str, Any]: } # 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": []} From c530768bc29b674826ab6bf9d803e3872ac37bf7 Mon Sep 17 00:00:00 2001 From: Jordan Leis Date: Fri, 19 Jun 2026 18:27:00 +0000 Subject: [PATCH 5/5] fix: resolve Copilot review findings in Cohere integration - cohere.py: replace time.sleep() with await asyncio.sleep() in both retry paths so backoff no longer blocks the event loop - modelSelection.ts: add Cohere branch in modelConfigToSelection() so auto-selected Cohere models get modelType "cohere" not "azure" - router.py: use COHERE_MODEL_NAME env var for model id instead of hard-coded "cohere-command-a" - test_chat_memory_service.py: add test asserting Cohere context window returns 256k tokens --- backend/api/server/router.py | 2 +- backend/services/llm/cohere.py | 4 ++-- backend/tests/test_chat_memory_service.py | 17 +++++++++++++++++ frontend/utils/modelSelection.ts | 5 ++++- 4 files changed, 24 insertions(+), 4 deletions(-) diff --git a/backend/api/server/router.py b/backend/api/server/router.py index 108e247..3183c97 100644 --- a/backend/api/server/router.py +++ b/backend/api/server/router.py @@ -580,7 +580,7 @@ async def get_available_models(): if os.getenv("COHERE_AZURE_ENDPOINT") and os.getenv("COHERE_AZURE_KEY"): models.append( { - "id": "cohere-command-a", + "id": os.getenv("COHERE_MODEL_NAME", "cohere-command-a"), "name": os.getenv("COHERE_DISPLAY_NAME", "Cohere Command A"), "provider": "Cohere", "model_type": "cohere", diff --git a/backend/services/llm/cohere.py b/backend/services/llm/cohere.py index 35a67d2..57131dd 100644 --- a/backend/services/llm/cohere.py +++ b/backend/services/llm/cohere.py @@ -150,7 +150,7 @@ async def _call_api( print( f"[CohereLLMClient] HTTP {resp.status_code} on attempt {attempt + 1}, retrying in {delay:.1f}s" ) - time.sleep(delay) + await asyncio.sleep(delay) continue # Non-retryable or last attempt @@ -169,7 +169,7 @@ async def _call_api( except requests.exceptions.Timeout: last_error = f"Cohere request timed out (attempt {attempt + 1})" if attempt < max_retries - 1: - time.sleep(base_delay * (2**attempt)) + await asyncio.sleep(base_delay * (2**attempt)) continue break except Exception as exc: 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/frontend/utils/modelSelection.ts b/frontend/utils/modelSelection.ts index e32cf5a..747b824 100644 --- a/frontend/utils/modelSelection.ts +++ b/frontend/utils/modelSelection.ts @@ -162,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,