diff --git a/backend/app/agents/base.py b/backend/app/agents/base.py index 6bdbd15..0d490bd 100644 --- a/backend/app/agents/base.py +++ b/backend/app/agents/base.py @@ -1,4 +1,4 @@ -"""Shared agentic loop runner used by all agents.""" +"""Shared agentic loop runner — supports Anthropic and OpenAI-compatible backends (Groq etc.).""" from __future__ import annotations import asyncio @@ -8,14 +8,25 @@ from dataclasses import dataclass, field from typing import Any -import anthropic - from .tools import TOOL_SCHEMAS, ToolContext, execute_tool logger = logging.getLogger(__name__) _RATE_LIMIT_WAITS = [15, 30, 60, 120] # seconds to wait on successive 429s +# OpenAI-compatible tool schema (used for Groq and any OpenAI endpoint) +_OPENAI_TOOL_SCHEMAS: list[dict] = [ + { + "type": "function", + "function": { + "name": t["name"], + "description": t.get("description", ""), + "parameters": t.get("input_schema", {"type": "object", "properties": {}}), + }, + } + for t in TOOL_SCHEMAS +] + @dataclass class AgentResult: @@ -25,8 +36,12 @@ class AgentResult: error: str | None = None +def _is_anthropic(client: Any) -> bool: + return type(client).__name__ == "AsyncAnthropic" + + async def run_agent( - client: anthropic.AsyncAnthropic, + client: Any, model: str, agent_name: str, system_prompt: str, @@ -34,43 +49,65 @@ async def run_agent( tool_context: ToolContext, max_turns: int = 12, ) -> AgentResult: + if _is_anthropic(client): + return await _run_anthropic(client, model, agent_name, system_prompt, initial_message, tool_context, max_turns) + return await _run_openai(client, model, agent_name, system_prompt, initial_message, tool_context, max_turns) + + +# --------------------------------------------------------------------------- +# Anthropic path +# --------------------------------------------------------------------------- + +async def _run_anthropic( + client: Any, + model: str, + agent_name: str, + system_prompt: str, + initial_message: str, + tool_context: ToolContext, + max_turns: int, +) -> AgentResult: + import anthropic as _anthropic + messages: list[dict] = [{"role": "user", "content": initial_message}] for turn in range(max_turns): response = None for attempt, wait in enumerate([0] + _RATE_LIMIT_WAITS): if wait: - logger.warning( - "Agent %s rate-limited on turn %d, waiting %ds (attempt %d)", - agent_name, turn, wait, attempt, - ) + logger.warning("Agent %s rate-limited turn %d, waiting %ds", agent_name, turn, wait) await asyncio.sleep(wait) try: response = await client.messages.create( model=model, - max_tokens=8192, - system=[ - { - "type": "text", - "text": system_prompt, - "cache_control": {"type": "ephemeral"}, - } - ], + max_tokens=2048, + system=[{"type": "text", "text": system_prompt, "cache_control": {"type": "ephemeral"}}], tools=TOOL_SCHEMAS, messages=messages, ) break - except anthropic.RateLimitError as exc: + except _anthropic.RateLimitError as exc: if attempt == len(_RATE_LIMIT_WAITS): - logger.error("Agent %s exhausted retries on turn %d", agent_name, turn) + logger.error("Agent %s exhausted retries turn %d", agent_name, turn) return AgentResult(name=agent_name, text="", error=str(exc)) except Exception as exc: - logger.error("Agent %s failed on turn %d: %s", agent_name, turn, exc) + logger.error("Agent %s failed turn %d: %s", agent_name, turn, exc) return AgentResult(name=agent_name, text="", error=str(exc)) + if response is None: return AgentResult(name=agent_name, text="", error="rate_limit_exhausted") - messages.append({"role": "assistant", "content": response.content}) + # Anthropic rejects messages where a TextBlock ends with trailing whitespace. + # Convert content blocks to dicts and strip to avoid the 400 error. + cleaned: list[dict] = [] + for blk in response.content: + if blk.type == "text": + cleaned.append({"type": "text", "text": blk.text.rstrip() or " "}) + elif blk.type == "tool_use": + cleaned.append({"type": "tool_use", "id": blk.id, "name": blk.name, "input": blk.input}) + else: + cleaned.append(blk) + messages.append({"role": "assistant", "content": cleaned}) if response.stop_reason == "end_turn": text = "".join(b.text for b in response.content if hasattr(b, "text")) @@ -86,19 +123,99 @@ async def run_agent( except Exception as exc: logger.warning("Tool %s raised: %s", block.name, exc) content = json.dumps({"error": str(exc)}) - tool_results.append( - { - "type": "tool_result", - "tool_use_id": block.id, - "content": content, - } - ) + tool_results.append({"type": "tool_result", "tool_use_id": block.id, "content": content}) messages.append({"role": "user", "content": tool_results}) logger.warning("Agent %s hit max_turns=%d", agent_name, max_turns) return AgentResult(name=agent_name, text="Max turns reached", error="max_turns_exceeded") +# --------------------------------------------------------------------------- +# OpenAI-compatible path (Groq, OpenAI, Ollama, etc.) +# --------------------------------------------------------------------------- + +async def _run_openai( + client: Any, + model: str, + agent_name: str, + system_prompt: str, + initial_message: str, + tool_context: ToolContext, + max_turns: int, +) -> AgentResult: + messages: list[dict] = [ + {"role": "system", "content": system_prompt}, + {"role": "user", "content": initial_message}, + ] + + for turn in range(max_turns): + response = None + for attempt, wait in enumerate([0] + _RATE_LIMIT_WAITS): + if wait: + logger.warning("Agent %s rate-limited turn %d, waiting %ds", agent_name, turn, wait) + await asyncio.sleep(wait) + try: + response = await client.chat.completions.create( + model=model, + max_tokens=2048, + tools=_OPENAI_TOOL_SCHEMAS, + messages=messages, + ) + break + except Exception as exc: + err = str(exc) + is_rate_limit = ( + "429" in err + or "rate_limit" in err.lower() + or "RateLimitError" in type(exc).__name__ + ) + if is_rate_limit and attempt < len(_RATE_LIMIT_WAITS): + continue + logger.error("Agent %s failed turn %d: %s", agent_name, turn, exc) + return AgentResult(name=agent_name, text="", error=err) + + if response is None: + return AgentResult(name=agent_name, text="", error="rate_limit_exhausted") + + choice = response.choices[0] + msg = choice.message + + if choice.finish_reason in ("stop", "end_turn", None) and not msg.tool_calls: + text = msg.content or "" + return AgentResult(name=agent_name, text=text, parsed=_extract_json(text)) + + if choice.finish_reason == "tool_calls" or msg.tool_calls: + # Append assistant message with tool_calls + messages.append({ + "role": "assistant", + "content": msg.content, + "tool_calls": [ + { + "id": tc.id, + "type": "function", + "function": {"name": tc.function.name, "arguments": tc.function.arguments}, + } + for tc in (msg.tool_calls or []) + ], + }) + for tc in (msg.tool_calls or []): + try: + args = json.loads(tc.function.arguments or "{}") + result = await execute_tool(tc.function.name, args, tool_context) + content = json.dumps(result, default=str) + except Exception as exc: + logger.warning("Tool %s raised: %s", tc.function.name, exc) + content = json.dumps({"error": str(exc)}) + messages.append({"role": "tool", "tool_call_id": tc.id, "content": content}) + else: + # Unexpected finish reason — treat as end + text = msg.content or "" + return AgentResult(name=agent_name, text=text, parsed=_extract_json(text)) + + logger.warning("Agent %s hit max_turns=%d", agent_name, max_turns) + return AgentResult(name=agent_name, text="Max turns reached", error="max_turns_exceeded") + + def _extract_json(text: str) -> dict[str, Any]: m = re.search(r"```json\s*(\{.*?\})\s*```", text, re.DOTALL) if m: diff --git a/backend/app/agents/bear_case_agent.py b/backend/app/agents/bear_case_agent.py index 12df9fe..e832309 100644 --- a/backend/app/agents/bear_case_agent.py +++ b/backend/app/agents/bear_case_agent.py @@ -5,9 +5,8 @@ from datetime import UTC, datetime from typing import Any -import anthropic -from app.core.config import get_settings +from app.core.config import get_active_model, get_settings from .base import AgentResult, run_agent from .tools import ToolContext @@ -51,7 +50,7 @@ async def run_bear_case_agent( sector_top_picks: list[str], sector_summaries: dict[str, Any], - client: anthropic.AsyncAnthropic, + client: object, tool_context: ToolContext, ) -> AgentResult: today = datetime.now(tz=UTC).strftime("%Y-%m-%d") @@ -76,7 +75,7 @@ async def run_bear_case_agent( return await run_agent( client=client, - model=get_settings().agent_model, + model=get_active_model(), agent_name="bear_case", system_prompt=_BEAR_SYSTEM, initial_message=initial_message, diff --git a/backend/app/agents/catalyst_agent.py b/backend/app/agents/catalyst_agent.py index b1825f0..d948099 100644 --- a/backend/app/agents/catalyst_agent.py +++ b/backend/app/agents/catalyst_agent.py @@ -4,9 +4,8 @@ import logging from datetime import UTC, datetime -import anthropic -from app.core.config import get_settings +from app.core.config import get_active_model, get_settings from .base import AgentResult, run_agent from .tools import ToolContext @@ -52,7 +51,7 @@ async def run_catalyst_agent( sector_top_picks: list[str], - client: anthropic.AsyncAnthropic, + client: object, tool_context: ToolContext, ) -> AgentResult: today = datetime.now(tz=UTC).strftime("%Y-%m-%d") @@ -72,7 +71,7 @@ async def run_catalyst_agent( return await run_agent( client=client, - model=get_settings().agent_model, + model=get_active_model(), agent_name="catalyst", system_prompt=_CATALYST_SYSTEM, initial_message=initial_message, diff --git a/backend/app/agents/daily_scan.py b/backend/app/agents/daily_scan.py index d7d3122..4ed700c 100644 --- a/backend/app/agents/daily_scan.py +++ b/backend/app/agents/daily_scan.py @@ -5,13 +5,12 @@ from datetime import UTC, datetime from typing import Any -import anthropic - from app.constants import REDIS_AGENT_ANALYSIS_KEY, REDIS_DAILY_SCAN_KEY from app.core.redis_client import cache_load_json, cache_save_json from app.db.session import async_session_factory from .base import run_agent +from .llm_client import make_agent_client from .tools import ToolContext logger = logging.getLogger(__name__) @@ -57,10 +56,10 @@ async def run_daily_scan() -> dict[str, Any]: if not active_picks: return {"skipped": True, "reason": "No active BUY picks to monitor"} - from app.core.config import get_settings - settings = get_settings() - if not settings.anthropic_api_key: - return {"skipped": True, "reason": "No ANTHROPIC_API_KEY configured"} + try: + client, sub_model, _ = make_agent_client() + except ValueError as exc: + return {"skipped": True, "reason": str(exc)} picks_summary = "\n".join( f"- {p['ticker']} ({p.get('horizon','?')}-term, {p.get('final_recommendation')}): " @@ -77,14 +76,11 @@ async def run_daily_scan() -> dict[str, Any]: "3. Output your JSON health check" ) - client = anthropic.AsyncAnthropic(api_key=settings.anthropic_api_key) - haiku_model = "claude-haiku-4-5-20251001" - async with async_session_factory() as session: tool_context = ToolContext(session=session, top_n=10) result = await run_agent( client=client, - model=haiku_model, + model=sub_model, agent_name="daily_scan", system_prompt=_SCAN_SYSTEM, initial_message=initial_message, diff --git a/backend/app/agents/debate_agents.py b/backend/app/agents/debate_agents.py index 1916fed..e24bb66 100644 --- a/backend/app/agents/debate_agents.py +++ b/backend/app/agents/debate_agents.py @@ -10,9 +10,8 @@ from datetime import UTC, datetime from typing import Any -import anthropic -from app.core.config import get_settings +from app.core.config import get_active_overseer_model, get_settings from .base import AgentResult, run_agent from .tools import ToolContext @@ -73,7 +72,7 @@ async def run_bull_debate_agent( ticker: str, sector_thesis: str, bear_objection: str, - client: anthropic.AsyncAnthropic, + client: object, tool_context: ToolContext, ) -> AgentResult: today = datetime.now(tz=UTC).strftime("%Y-%m-%d") @@ -86,7 +85,7 @@ async def run_bull_debate_agent( ) return await run_agent( client=client, - model=get_settings().agent_model, + model=get_active_overseer_model(), agent_name=f"bull_debate_{ticker}", system_prompt=_BULL_SYSTEM, initial_message=message, @@ -99,7 +98,7 @@ async def run_bear_rebuttal_agent( ticker: str, bear_thesis: str, sector_objection: str, - client: anthropic.AsyncAnthropic, + client: object, tool_context: ToolContext, ) -> AgentResult: today = datetime.now(tz=UTC).strftime("%Y-%m-%d") @@ -112,7 +111,7 @@ async def run_bear_rebuttal_agent( ) return await run_agent( client=client, - model=get_settings().agent_model, + model=get_active_overseer_model(), agent_name=f"bear_rebuttal_{ticker}", system_prompt=_BEAR_REBUTTAL_SYSTEM, initial_message=message, @@ -125,7 +124,7 @@ async def run_debate_round( overseer_parsed: dict[str, Any], bear_parsed: dict[str, Any], sector_summaries: dict[str, str], - client: anthropic.AsyncAnthropic, + client: object, tool_context: ToolContext, ) -> dict[str, Any]: """Run debate agents on STRONG_BUY and AVOID tickers in parallel.""" diff --git a/backend/app/agents/llm_client.py b/backend/app/agents/llm_client.py new file mode 100644 index 0000000..8e49d17 --- /dev/null +++ b/backend/app/agents/llm_client.py @@ -0,0 +1,107 @@ +"""LLM client factory — returns the right async client based on llm_provider config.""" +from __future__ import annotations + +from typing import Any + +from app.core.config import get_settings + + +def make_agent_client() -> tuple[Any, str, str]: + """Return (client, model, overseer_model) based on the configured provider. + + Raises ValueError if the required API key is missing. + """ + settings = get_settings() + provider = settings.llm_provider.lower() + + if provider == "cerebras": + if not settings.cerebras_api_key: + raise ValueError( + "llm_provider=cerebras but CEREBRAS_API_KEY is not set. " + "Get a free key at https://cloud.cerebras.ai and add CEREBRAS_API_KEY to your .env" + ) + from openai import AsyncOpenAI + client = AsyncOpenAI( + api_key=settings.cerebras_api_key, + base_url="https://api.cerebras.ai/v1", + ) + return client, settings.cerebras_model, settings.cerebras_overseer_model + + if provider == "gemini": + if not settings.gemini_api_key: + raise ValueError( + "llm_provider=gemini but GEMINI_API_KEY is not set. " + "Get a free key at https://aistudio.google.com and add GEMINI_API_KEY to your .env" + ) + from openai import AsyncOpenAI + client = AsyncOpenAI( + api_key=settings.gemini_api_key, + base_url="https://generativelanguage.googleapis.com/v1beta/openai/", + ) + return client, settings.gemini_model, settings.gemini_overseer_model + + if provider == "groq": + if not settings.groq_api_key: + raise ValueError( + "llm_provider=groq but GROQ_API_KEY is not set. " + "Sign up free at https://console.groq.com and add GROQ_API_KEY to your .env" + ) + from openai import AsyncOpenAI + client = AsyncOpenAI( + api_key=settings.groq_api_key, + base_url="https://api.groq.com/openai/v1", + ) + return client, settings.groq_model, settings.groq_overseer_model + + if provider == "anthropic": + if not settings.anthropic_api_key: + raise ValueError("llm_provider=anthropic but ANTHROPIC_API_KEY is not set.") + import anthropic + client = anthropic.AsyncAnthropic(api_key=settings.anthropic_api_key) + return client, settings.agent_model, settings.agent_overseer_model + + raise ValueError(f"Unknown llm_provider={provider!r}. Must be 'gemini', 'groq' or 'anthropic'.") + + +def make_overseer_client() -> tuple[Any, str]: + """Return (client, model) for the overseer + debate phase. + + Uses OVERSEER_LLM_PROVIDER if set, otherwise falls back to the main llm_provider. + """ + settings = get_settings() + provider = (settings.overseer_llm_provider or settings.llm_provider).lower() + + if provider == "anthropic": + if not settings.anthropic_api_key: + raise ValueError( + "overseer_llm_provider=anthropic but ANTHROPIC_API_KEY is not set." + ) + import anthropic + client = anthropic.AsyncAnthropic(api_key=settings.anthropic_api_key) + return client, settings.agent_overseer_model + + if provider == "cerebras": + from openai import AsyncOpenAI + client = AsyncOpenAI( + api_key=settings.cerebras_api_key, + base_url="https://api.cerebras.ai/v1", + ) + return client, settings.cerebras_overseer_model + + if provider == "gemini": + from openai import AsyncOpenAI + client = AsyncOpenAI( + api_key=settings.gemini_api_key, + base_url="https://generativelanguage.googleapis.com/v1beta/openai/", + ) + return client, settings.gemini_overseer_model + + if provider == "groq": + from openai import AsyncOpenAI + client = AsyncOpenAI( + api_key=settings.groq_api_key, + base_url="https://api.groq.com/openai/v1", + ) + return client, settings.groq_overseer_model + + raise ValueError(f"Unknown overseer_llm_provider={provider!r}.") diff --git a/backend/app/agents/overseer.py b/backend/app/agents/overseer.py index 3225acc..496ddf5 100644 --- a/backend/app/agents/overseer.py +++ b/backend/app/agents/overseer.py @@ -5,9 +5,8 @@ from datetime import UTC, datetime from typing import Any -import anthropic -from app.core.config import get_settings +from app.core.config import get_active_overseer_model, get_settings from .base import AgentResult, run_agent from .tools import ToolContext @@ -96,7 +95,7 @@ async def run_overseer( sub_results: list[AgentResult], catalyst_result: AgentResult | None, bear_result: AgentResult | None, - client: anthropic.AsyncAnthropic, + client: object, tool_context: ToolContext, debate_context: dict[str, Any] | None = None, risk_context: dict[str, Any] | None = None, @@ -173,12 +172,60 @@ async def run_overseer( "6. Output the final portfolio JSON") ) - return await run_agent( + result = await run_agent( client=client, - model=get_settings().agent_overseer_model, + model=get_active_overseer_model(), agent_name="overseer", system_prompt=_OVERSEER_SYSTEM, initial_message=initial_message, tool_context=tool_context, max_turns=15, ) + + # If the model produced analysis but no JSON, do a direct one-shot extraction call (no tools). + if result.error is None and not result.parsed and result.text: + logger.info("Overseer produced no JSON — running direct extraction call") + try: + from .base import _extract_json, _is_anthropic + + extraction_prompt = ( + "You are a JSON formatter. Convert the portfolio analysis below into the required JSON schema.\n" + "Output ONLY valid JSON starting with { — no markdown fences, no preamble, no explanation.\n\n" + "ANALYSIS:\n" + f"{result.text}\n\n" + "REQUIRED SCHEMA:\n" + '{"market_overview":"...","portfolio_thesis":"...","verified_trades":[{"ticker":"...",' + '"asset_class":"stock","sector":"...","ml_signal":"BUY|HOLD",' + '"final_recommendation":"STRONG_BUY|BUY|HOLD|AVOID","conviction":"high|medium|low",' + '"horizon":"short|medium","position_size_pct":5,"agent_consensus":"agree",' + '"catalyst":null,"catalyst_date":null,"supporting_themes":[],"risk_factors":[],' + '"what_breaks_thesis":"...","suggested_action":"..."}],' + '"watchlist":[],"risk_notes":"..."}' + ) + + model = get_active_overseer_model() + if _is_anthropic(client): + resp = await client.messages.create( + model=model, + max_tokens=4096, + messages=[{"role": "user", "content": extraction_prompt}], + ) + raw = "".join(b.text for b in resp.content if hasattr(b, "text")) + else: + resp = await client.chat.completions.create( + model=model, + max_tokens=4096, + messages=[{"role": "user", "content": extraction_prompt}], + ) + raw = resp.choices[0].message.content or "" + + parsed = _extract_json(raw) + if parsed: + logger.info("JSON extraction succeeded — %d trades", len(parsed.get("verified_trades", []))) + return AgentResult(name="overseer", text=raw, parsed=parsed) + else: + logger.warning("JSON extraction follow-up also returned no parsed JSON; raw=%r", raw[:300]) + except Exception as exc: + logger.error("JSON extraction follow-up failed: %s", exc) + + return result diff --git a/backend/app/agents/pipeline.py b/backend/app/agents/pipeline.py index c7d9097..06fd40e 100644 --- a/backend/app/agents/pipeline.py +++ b/backend/app/agents/pipeline.py @@ -7,7 +7,6 @@ from datetime import UTC, datetime from typing import Any -import anthropic from sqlalchemy.ext.asyncio import AsyncSession from app.constants import REDIS_AGENT_ANALYSIS_KEY, REDIS_AGENT_META_KEY, REDIS_PORTFOLIO_RISK_KEY @@ -20,6 +19,7 @@ from .bear_case_agent import run_bear_case_agent from .catalyst_agent import run_catalyst_agent from .debate_agents import run_debate_round +from .llm_client import make_agent_client, make_overseer_client from .memory import save_agent_memories from .overseer import run_overseer from .sub_agents import ALL_SUB_AGENTS, run_all_sub_agents @@ -47,12 +47,11 @@ def _collect_summaries(sub_results: list) -> dict[str, str]: async def run_agent_pipeline(session: AsyncSession) -> dict[str, Any]: settings = get_settings() - if not settings.anthropic_api_key: - raise ValueError("ANTHROPIC_API_KEY is not configured.") + client, sub_model, overseer_model = make_agent_client() + o_client, o_model = make_overseer_client() run_id = str(uuid.uuid4()) - client = anthropic.AsyncAnthropic(api_key=settings.anthropic_api_key) - tool_context = ToolContext(session=session, top_n=settings.agent_top_n_per_sector) + tool_context = ToolContext(session_factory=async_session_factory, top_n=settings.agent_top_n_per_sector) # Phase 1 — sector agents in parallel (semaphore-limited) logger.info("Pipeline %s: running %d sector agents", run_id, len(ALL_SUB_AGENTS)) @@ -89,7 +88,7 @@ async def run_agent_pipeline(session: AsyncSession) -> dict[str, Any]: # Phase 3 — overseer (initial synthesis) logger.info("Phase 3: overseer initial synthesis (calibration for %d agents)", len(calibration)) overseer_result = await run_overseer( - sub_results, catalyst_result, bear_result, client, tool_context, + sub_results, catalyst_result, bear_result, o_client, tool_context, risk_context=risk_context, agent_calibration=calibration, ) @@ -115,14 +114,14 @@ async def run_agent_pipeline(session: AsyncSession) -> dict[str, Any]: overseer_result.parsed, bear_result.parsed if bear_result and not bear_result.error else {}, summaries, - client, + o_client, tool_context, ) # Re-run overseer with debate context if debate produced results if debate_results["bull_debates"] or debate_results["bear_rebuttals"]: logger.info("Phase 4b: overseer final synthesis with debate context") overseer_result = await run_overseer( - sub_results, catalyst_result, bear_result, client, tool_context, + sub_results, catalyst_result, bear_result, o_client, tool_context, debate_context=debate_results, risk_context=risk_context, agent_calibration=calibration, diff --git a/backend/app/agents/sub_agents.py b/backend/app/agents/sub_agents.py index 362994d..ee746ff 100644 --- a/backend/app/agents/sub_agents.py +++ b/backend/app/agents/sub_agents.py @@ -6,9 +6,8 @@ from dataclasses import dataclass from datetime import UTC, datetime -import anthropic -from app.core.config import get_settings +from app.core.config import get_active_model, get_settings from .base import AgentResult, run_agent from .tools import ToolContext @@ -228,12 +227,15 @@ def _spec(name: str, role: str, coverage: str, focus: str) -> SubAgentSpec: ] -_CONCURRENCY = 3 # max simultaneous Anthropic API calls (Tier-1 rate limit safety) +def _concurrency() -> int: + from app.core.config import get_settings + # Groq free tier: 30 RPM, 6000 TPM — keep lower concurrency to avoid token-limit bursts + return 2 if get_settings().llm_provider == "groq" else 3 async def run_sub_agent( spec: SubAgentSpec, - client: anthropic.AsyncAnthropic, + client: object, tool_context: ToolContext, sem: asyncio.Semaphore, ) -> AgentResult: @@ -244,7 +246,7 @@ async def run_sub_agent( async with sem: return await run_agent( client=client, - model=get_settings().agent_model, + model=get_active_model(), agent_name=spec.name, system_prompt=spec.system_prompt, initial_message=initial_message, @@ -254,10 +256,10 @@ async def run_sub_agent( async def run_all_sub_agents( - client: anthropic.AsyncAnthropic, + client: object, tool_context: ToolContext, ) -> list[AgentResult]: - sem = asyncio.Semaphore(_CONCURRENCY) + sem = asyncio.Semaphore(_concurrency()) tasks = [run_sub_agent(spec, client, tool_context, sem) for spec in ALL_SUB_AGENTS] raw = await asyncio.gather(*tasks, return_exceptions=True) results: list[AgentResult] = [] diff --git a/backend/app/agents/tools.py b/backend/app/agents/tools.py index 318d3e3..35e9465 100644 --- a/backend/app/agents/tools.py +++ b/backend/app/agents/tools.py @@ -7,6 +7,8 @@ from datetime import date, timedelta from typing import Any +from typing import Callable + from sqlalchemy import desc, func, select from sqlalchemy.ext.asyncio import AsyncSession @@ -268,7 +270,7 @@ @dataclass class ToolContext: - session: AsyncSession + session_factory: Callable top_n: int = 10 @@ -285,8 +287,9 @@ async def _web_search(query: str, max_results: int = 5) -> list[dict]: from tavily import TavilyClient # type: ignore[import] client = TavilyClient(api_key=settings.tavily_api_key) - response = await asyncio.to_thread( - client.search, query, max_results=max_results, search_depth="advanced" + response = await asyncio.wait_for( + asyncio.to_thread(client.search, query, max_results=max_results, search_depth="basic"), + timeout=15.0, ) results = response.get("results", []) return [ @@ -298,6 +301,9 @@ async def _web_search(query: str, max_results: int = 5) -> list[dict]: } for r in results ] + except asyncio.TimeoutError: + logger.warning("Tavily search timed out for %r", query) + return [{"title": "Search timeout", "content": "Web search timed out. Analysis based on database signals only."}] except Exception as exc: logger.warning("Tavily search failed for %r: %s", query, exc) return [{"title": "Search error", "content": str(exc)}] @@ -741,16 +747,6 @@ async def execute_tool(name: str, inputs: dict, ctx: ToolContext) -> Any: return await _web_search(inputs.get("query", ""), inputs.get("max_results", 5)) if name == "get_commodity_signals": return await _get_commodity_signals() - if name == "get_stock_rankings": - return await _get_stock_rankings( - inputs.get("sector_group", ""), inputs.get("top_n", ctx.top_n), ctx.session - ) - if name == "get_macro_indicators": - return await _get_macro_indicators(ctx.session) - if name == "get_price_history": - return await _get_price_history(inputs.get("ticker", ""), inputs.get("days", 30), ctx.session) - if name == "get_sentiment_scores": - return await _get_sentiment_scores(inputs.get("tickers", []), ctx.session) if name == "get_fundamentals": return await _get_fundamentals(inputs.get("ticker", "")) if name == "get_earnings_calendar": @@ -761,10 +757,24 @@ async def execute_tool(name: str, inputs: dict, ctx: ToolContext) -> Any: return await _get_options_context(inputs.get("ticker", "")) if name == "get_insider_activity": return await _get_insider_activity(inputs.get("ticker", "")) - if name == "get_cot_positioning": - return await _get_cot_positioning(inputs.get("ticker", ""), inputs.get("weeks", 26), ctx.session) - if name == "search_memory": - return await _search_memory(inputs.get("query", ""), inputs.get("agent_name"), ctx.session) - if name == "get_economic_calendar": - return await _get_economic_calendar(inputs.get("days", 30), ctx.session) + + # DB tools — each gets its own session to allow safe concurrent execution + async with ctx.session_factory() as session: + if name == "get_stock_rankings": + return await _get_stock_rankings( + inputs.get("sector_group", ""), inputs.get("top_n", ctx.top_n), session + ) + if name == "get_macro_indicators": + return await _get_macro_indicators(session) + if name == "get_price_history": + return await _get_price_history(inputs.get("ticker", ""), inputs.get("days", 30), session) + if name == "get_sentiment_scores": + return await _get_sentiment_scores(inputs.get("tickers", []), session) + if name == "get_cot_positioning": + return await _get_cot_positioning(inputs.get("ticker", ""), inputs.get("weeks", 26), session) + if name == "search_memory": + return await _search_memory(inputs.get("query", ""), inputs.get("agent_name"), session) + if name == "get_economic_calendar": + return await _get_economic_calendar(inputs.get("days", 30), session) + return {"error": f"Unknown tool: {name}"} diff --git a/backend/app/core/config.py b/backend/app/core/config.py index 5e4d1b1..9668f92 100644 --- a/backend/app/core/config.py +++ b/backend/app/core/config.py @@ -24,13 +24,40 @@ class Settings(BaseSettings): # Set in .env as ``API_KEY=`` for any deployed instance. api_key: str = "" - # Agent pipeline (Claude + web search). + # Agent pipeline — LLM provider selection. + # "groq" = free tier with Llama 3.3 70B (sign up at console.groq.com, no billing required) + # "anthropic" = Claude models (paid — use as fallback or for higher quality) + llm_provider: str = "cerebras" # "cerebras" | "gemini" | "groq" | "anthropic" + + # Cerebras settings (free — 1M tokens/day, sign up at cloud.cerebras.ai) + cerebras_api_key: str = "" + cerebras_model: str = "zai-glm-4.7" + cerebras_overseer_model: str = "zai-glm-4.7" + + # Gemini settings (free — 20 req/day on 2.5-flash) + gemini_api_key: str = "" + gemini_model: str = "gemini-2.5-flash" + gemini_overseer_model: str = "gemini-2.5-flash" + + # Groq settings (free — 100K TPD on 70B) + groq_api_key: str = "" + groq_model: str = "llama-3.3-70b-versatile" + groq_overseer_model: str = "llama-3.3-70b-versatile" + + # Anthropic settings (fallback / paid) anthropic_api_key: str = "" + agent_model: str = "claude-haiku-4-5-20251001" + agent_overseer_model: str = "claude-haiku-4-5-20251001" + + # Optionally use a different provider just for the overseer + debate agents. + # Empty string = use llm_provider for everything. + # Example: set OVERSEER_LLM_PROVIDER=anthropic to use Claude for the overseer + # while sub-agents run on the cheaper/free llm_provider. + overseer_llm_provider: str = "" + tavily_api_key: str = "" agent_top_n_per_sector: int = 10 - agent_model: str = "claude-sonnet-4-6" - agent_overseer_model: str = "claude-sonnet-4-6" - # Set to true to run the full agent pipeline (LLM costs) on schedule (Mon-Fri 09:00 NY). + # Set to true to run the full agent pipeline on schedule (Mon-Fri 09:00 NY). agent_analysis_scheduled: bool = False @@ -39,6 +66,34 @@ def get_settings() -> Settings: return Settings() +def get_active_model(*, overseer: bool = False) -> str: + """Return the correct model name for the configured llm_provider.""" + s = get_settings() + if s.llm_provider == "cerebras": + return s.cerebras_overseer_model if overseer else s.cerebras_model + if s.llm_provider == "gemini": + return s.gemini_overseer_model if overseer else s.gemini_model + if s.llm_provider == "groq": + return s.groq_overseer_model if overseer else s.groq_model + return s.agent_overseer_model if overseer else s.agent_model + + +def get_active_overseer_model() -> str: + """Return the model name for the overseer/debate phase. + + Respects overseer_llm_provider if set, otherwise falls back to get_active_model(overseer=True). + """ + s = get_settings() + provider = s.overseer_llm_provider or s.llm_provider + if provider == "cerebras": + return s.cerebras_overseer_model + if provider == "gemini": + return s.gemini_overseer_model + if provider == "groq": + return s.groq_overseer_model + return s.agent_overseer_model + + def parse_cors(s: str) -> list[str]: if s.strip() == "*": return ["*"] diff --git a/backend/app/scheduler.py b/backend/app/scheduler.py index 2415d5e..bd5f219 100644 --- a/backend/app/scheduler.py +++ b/backend/app/scheduler.py @@ -423,7 +423,7 @@ async def daily_agent_analysis() -> None: misfire_grace_time=7200, ) - if settings.agent_analysis_scheduled and settings.anthropic_api_key: + if settings.agent_analysis_scheduled and (settings.cerebras_api_key or settings.anthropic_api_key or settings.gemini_api_key or settings.groq_api_key): sched.add_job( daily_agent_analysis, CronTrigger(day_of_week="mon-fri", hour=9, minute=0), diff --git a/backend/requirements.txt b/backend/requirements.txt index d723635..ef9aeed 100644 --- a/backend/requirements.txt +++ b/backend/requirements.txt @@ -33,5 +33,6 @@ lxml==5.3.0 pytest==8.3.4 pytest-asyncio==0.25.0 anthropic>=0.40.0 +openai>=1.50.0 tavily-python>=0.3.3 sentence-transformers>=3.0.0 diff --git a/backend/run_overseer_only.py b/backend/run_overseer_only.py new file mode 100644 index 0000000..5909e5e --- /dev/null +++ b/backend/run_overseer_only.py @@ -0,0 +1,141 @@ +"""One-shot script: run just the overseer phase using cached sub-agent results.""" +import asyncio +import logging +from datetime import UTC, datetime + +logging.basicConfig(level=logging.INFO, format="%(levelname)s %(name)s: %(message)s") +logger = logging.getLogger("run_overseer_only") + + +async def main() -> None: + from app.agents.base import AgentResult + from app.agents.debate_agents import run_debate_round + from app.agents.llm_client import make_overseer_client + from app.agents.memory import save_agent_memories + from app.agents.overseer import run_overseer + from app.constants import REDIS_AGENT_ANALYSIS_KEY, REDIS_AGENT_META_KEY, REDIS_PORTFOLIO_RISK_KEY + from app.core.redis_client import cache_load_json, cache_save_json + from app.db.session import async_session_factory + from app.services.recommendations_service import get_agent_calibration, save_recommendations + from app.agents.tools import ToolContext + + # Load cached pipeline data + logger.info("Loading cached sub-agent results from Redis...") + data = await cache_load_json(REDIS_AGENT_ANALYSIS_KEY) + if not data: + logger.error("No cached analysis found. Run the full pipeline first.") + return + + run_id = data["run_id"] + logger.info("Using run_id=%s", run_id) + + # Reconstruct AgentResult objects + sub_results = [ + AgentResult(name=r["name"], text=r.get("text", ""), parsed=r.get("parsed", {}), error=r.get("error")) + for r in data.get("sub_reports", []) + ] + cr = data.get("catalyst_report") or {} + br = data.get("bear_report") or {} + catalyst_result = AgentResult(name="catalyst", text=cr.get("text", ""), parsed=cr.get("parsed", {}), error=cr.get("error")) if cr else None + bear_result = AgentResult(name="bear_case", text=br.get("text", ""), parsed=br.get("parsed", {}), error=br.get("error")) if br else None + + logger.info( + "Loaded: %d sub-agents, catalyst=%s, bear=%s", + len(sub_results), + "ok" if catalyst_result and not catalyst_result.error else "missing", + "ok" if bear_result and not bear_result.error else "missing", + ) + + # Build tool context and load context data + tool_context = ToolContext(session_factory=async_session_factory, top_n=10) + risk_context = await cache_load_json(REDIS_PORTFOLIO_RISK_KEY) + + try: + async with async_session_factory() as cal_session: + calibration = await get_agent_calibration(cal_session) + except Exception as exc: + logger.warning("Calibration load failed: %s", exc) + calibration = {} + + # Create overseer client (Anthropic) + o_client, o_model = make_overseer_client() + logger.info("Overseer client: %s, model: %s", type(o_client).__name__, o_model) + + # Phase 3 — overseer initial synthesis + logger.info("Running overseer initial synthesis...") + overseer_result = await run_overseer( + sub_results, catalyst_result, bear_result, o_client, tool_context, + risk_context=risk_context, + agent_calibration=calibration, + ) + logger.info("Overseer done. error=%s, trades=%d", + overseer_result.error, + len(overseer_result.parsed.get("verified_trades", [])), + ) + + # Phase 4 — debate round + summaries = {r.name: r.parsed.get("summary", r.text[:200]) for r in sub_results if not r.error} + debate_results: dict = {"bull_debates": {}, "bear_rebuttals": {}} + + if overseer_result.error is None and overseer_result.parsed.get("verified_trades"): + strong_buys = [t["ticker"] for t in overseer_result.parsed["verified_trades"] if t.get("final_recommendation") == "STRONG_BUY"] + avoids = [t["ticker"] for t in overseer_result.parsed["verified_trades"] if t.get("final_recommendation") == "AVOID"] + + if strong_buys or avoids: + logger.info("Running debate round — %d STRONG_BUY, %d AVOID", len(strong_buys), len(avoids)) + try: + debate_results = await run_debate_round( + overseer_result.parsed, + bear_result.parsed if bear_result and not bear_result.error else {}, + summaries, + o_client, + tool_context, + ) + if debate_results["bull_debates"] or debate_results["bear_rebuttals"]: + logger.info("Re-running overseer with debate context...") + overseer_result = await run_overseer( + sub_results, catalyst_result, bear_result, o_client, tool_context, + debate_context=debate_results, + risk_context=risk_context, + agent_calibration=calibration, + ) + except Exception as exc: + logger.error("Debate round failed: %s", exc) + + # Save results back to Redis + overseer_ok = overseer_result.error is None and bool(overseer_result.parsed.get("verified_trades")) + verified_trades = overseer_result.parsed.get("verified_trades", []) if overseer_ok else [] + logger.info("Overseer ok=%s, verified_trades=%d", overseer_ok, len(verified_trades)) + + updated_data = {**data} + updated_data["overseer"] = {"text": overseer_result.text, "parsed": overseer_result.parsed, "error": overseer_result.error} + updated_data["debate_report"] = debate_results + updated_data["generated_at"] = datetime.now(tz=UTC).isoformat() + + await cache_save_json(REDIS_AGENT_ANALYSIS_KEY, updated_data) + + meta = await cache_load_json(REDIS_AGENT_META_KEY) or {} + meta.update({ + "overseer_ok": overseer_ok, + "debate_tickers": list(debate_results.get("bull_debates", {}).keys()), + "generated_at": updated_data["generated_at"], + }) + await cache_save_json(REDIS_AGENT_META_KEY, meta) + + # Save recommendations to DB if we got trades + if verified_trades: + try: + async with async_session_factory() as rec_session: + await save_recommendations(run_id, verified_trades, rec_session) + logger.info("Saved %d recommendations to DB", len(verified_trades)) + except Exception as exc: + logger.error("Failed to save recommendations: %s", exc) + + logger.info("Done. Overseer result saved to Redis.") + if verified_trades: + for t in verified_trades: + logger.info(" %s — %s", t.get("ticker"), t.get("final_recommendation")) + + +if __name__ == "__main__": + asyncio.run(main())