File size: 13,624 Bytes
69e310f
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
274
275
276
277
278
279
280
281
282
283
284
285
286
287
288
289
290
291
292
293
294
295
296
297
298
299
300
301
302
303
304
305
306
307
308
309
310
311
312
313
314
315
316
317
318
319
320
321
322
323
324
325
326
327
328
329
330
331
332
333
334
335
336
337
338
339
340
341
342
343
344
345
346
347
348
349
350
351
352
353
354
355
356
357
358
359
360
361
362
363
364
365
366
367
368
369
370
"""MCP client the agents use. Tools are consumed over MCP, never called directly.

Two transports, same protocol and same tool schemas:

``stdio`` (default)
    Spawns ``python -m app.mcp_server`` as a child process and speaks JSON-RPC
    over its stdin/stdout — exactly how any other MCP client would connect.

``inmemory``
    Runs the same :class:`MCPServer` object over in-memory streams. Still the
    real MCP protocol; used by tests and the eval harness to avoid paying process
    spawn cost 20 times over.

The client re-enforces the tool whitelist. Even a compromised agent cannot ask
for a tool the server does not expose, and cannot reach anything but these four.
"""

from __future__ import annotations

import json
import logging
import os
import sys
import time
from collections.abc import Awaitable, Callable, Iterator
from contextlib import contextmanager
from dataclasses import dataclass, field
from pathlib import Path
from types import TracebackType
from typing import Any, Self

from mcp import Client, StdioServerParameters, stdio_client

from app.core.settings import Settings, get_settings
from app.mcp_server.providers import ProviderContext
from app.mcp_server.registry import TOOL_NAMES, build_server
from app.models.market import Fundamentals, Metrics, NewsFeed, PriceBar, PriceHistory

logger = logging.getLogger(__name__)

#: apps/api — the directory that must be importable for `-m app.mcp_server`.
PACKAGE_ROOT = Path(__file__).resolve().parents[2]


class McpToolError(RuntimeError):
    """A tool call could not be completed at the protocol level."""


_shared_server: Any | None = None


def shared_inmemory_server(settings: Settings) -> Any:
    """One process-wide in-memory MCP server, so its provider cache is shared.

    Opt-in via ``MCP_SHARED_SERVER``. The eval harness uses it so twenty runs
    make a handful of provider calls rather than a hundred — the graph, the MCP
    protocol and the verifier are all still exercised in full.
    """
    global _shared_server  # noqa: PLW0603 - deliberate process-wide singleton
    if _shared_server is None:
        _shared_server = build_server(
            ProviderContext(
                min_interval_seconds=settings.provider_min_interval_seconds,
                max_attempts=settings.provider_max_attempts,
                retry_backoff_seconds=settings.provider_retry_backoff_seconds,
            )
        )
    return _shared_server


def reset_shared_server() -> None:
    """Test hook: drop the shared server and its cache."""
    global _shared_server  # noqa: PLW0603 - deliberate process-wide singleton
    _shared_server = None


@dataclass(frozen=True, slots=True)
class ToolCallRecord:
    """Telemetry row for one MCP tool call — args and ms timing, as the UI shows."""

    tool: str
    arguments: dict[str, Any]
    duration_ms: float
    ok: bool
    summary: str
    error: str | None = None

    def to_dict(self) -> dict[str, Any]:
        return {
            "tool": self.tool,
            "arguments": self.arguments,
            "duration_ms": round(self.duration_ms, 1),
            "ok": self.ok,
            "summary": self.summary,
            "error": self.error,
        }


@dataclass(slots=True)
class _Args:
    """Compact, log-safe rendering of tool arguments for telemetry."""

    raw: dict[str, Any]

    def compact(self) -> dict[str, Any]:
        out: dict[str, Any] = {}
        for key, value in self.raw.items():
            if isinstance(value, list):
                out[key] = f"[{len(value)} items]"
            elif isinstance(value, str) and len(value) > 80:
                out[key] = value[:77] + "…"
            else:
                out[key] = value
        return out


@dataclass(slots=True)
class McpToolClient:
    """Async context manager wrapping one MCP session for one run.

    `emitter` lets the run publish each tool call to the live telemetry feed the
    instant it completes, so the UI shows calls scrolling rather than arriving in
    batches. `collect()` gives a node its own record slice without racing the
    other agent running in the same superstep.
    """

    settings: Settings = field(default_factory=get_settings)
    records: list[ToolCallRecord] = field(default_factory=list)
    emitter: Callable[[ToolCallRecord], Awaitable[None]] | None = None
    _sinks: list[list[ToolCallRecord]] = field(default_factory=list)
    _client: Client | None = None
    _stack: Any = None

    @contextmanager
    def collect(self) -> Iterator[list[ToolCallRecord]]:
        """Capture only the tool calls made inside this block."""
        sink: list[ToolCallRecord] = []
        self._sinks.append(sink)
        try:
            yield sink
        finally:
            if sink in self._sinks:
                self._sinks.remove(sink)

    async def _record(self, record: ToolCallRecord) -> None:
        self.records.append(record)
        for sink in self._sinks:
            sink.append(record)
        if self.emitter is not None:
            await self.emitter(record)

    async def __aenter__(self) -> Self:
        from contextlib import AsyncExitStack

        stack = AsyncExitStack()
        self._stack = stack
        transport = self.settings.mcp_transport
        try:
            if transport == "inmemory":
                server = (
                    shared_inmemory_server(self.settings)
                    if self.settings.mcp_shared_server
                    else build_server(
                        ProviderContext(
                            min_interval_seconds=self.settings.provider_min_interval_seconds
                        )
                    )
                )
                client = Client(server, read_timeout_seconds=self.settings.mcp_call_timeout_seconds)
            else:
                env = dict(os.environ)
                existing = env.get("PYTHONPATH", "")
                env["PYTHONPATH"] = (
                    f"{PACKAGE_ROOT}{os.pathsep}{existing}" if existing else str(PACKAGE_ROOT)
                )
                env["PROVIDER_MIN_INTERVAL_SECONDS"] = str(
                    self.settings.provider_min_interval_seconds
                )
                params = StdioServerParameters(
                    command=sys.executable,
                    args=["-m", "app.mcp_server"],
                    env=env,
                    cwd=str(PACKAGE_ROOT),
                )
                client = Client(
                    stdio_client(params),
                    read_timeout_seconds=self.settings.mcp_call_timeout_seconds,
                )
            self._client = await stack.enter_async_context(client)
        except Exception:
            await stack.aclose()
            self._stack = None
            raise
        return self

    async def __aexit__(
        self,
        exc_type: type[BaseException] | None,
        exc: BaseException | None,
        tb: TracebackType | None,
    ) -> None:
        stack, self._stack = self._stack, None
        self._client = None
        if stack is not None:
            await stack.aclose()

    # ------------------------------------------------------------------ raw ---
    async def list_tool_names(self) -> list[str]:
        """Names advertised by the connected server (discoverability check)."""
        client = self._require_client()
        listing = await client.list_tools()
        return sorted(tool.name for tool in listing.tools)

    async def list_tool_specs(self) -> list[dict[str, Any]]:
        """Full advertised tool documentation, exactly as any MCP client sees it."""
        client = self._require_client()
        listing = await client.list_tools()
        specs: list[dict[str, Any]] = []
        for tool in sorted(listing.tools, key=lambda t: t.name):
            schema = getattr(tool, "inputSchema", None) or getattr(tool, "input_schema", None) or {}
            specs.append(
                {
                    "name": tool.name,
                    "description": (tool.description or "").strip(),
                    "input_schema": dict(schema),
                }
            )
        return specs

    async def call(self, name: str, arguments: dict[str, Any]) -> dict[str, Any]:
        """Call one whitelisted MCP tool and return its structured payload."""
        if name not in TOOL_NAMES:
            raise McpToolError(f"tool '{name}' is not on the AlphaBrief MCP whitelist")
        client = self._require_client()
        started = time.perf_counter()
        try:
            result = await client.call_tool(name, arguments)
        except Exception as exc:
            elapsed = (time.perf_counter() - started) * 1000
            await self._record(
                ToolCallRecord(
                    tool=name,
                    arguments=_Args(arguments).compact(),
                    duration_ms=elapsed,
                    ok=False,
                    summary="transport error",
                    error=f"{type(exc).__name__}: {exc}",
                )
            )
            raise McpToolError(f"MCP call to '{name}' failed: {exc}") from exc

        elapsed = (time.perf_counter() - started) * 1000
        try:
            payload = _extract_payload(result)
        except McpToolError as exc:
            await self._record(
                ToolCallRecord(
                    tool=name,
                    arguments=_Args(arguments).compact(),
                    duration_ms=elapsed,
                    ok=False,
                    summary="tool error",
                    error=str(exc),
                )
            )
            raise

        tool_error = payload.get("error") if isinstance(payload, dict) else None
        await self._record(
            ToolCallRecord(
                tool=name,
                arguments=_Args(arguments).compact(),
                duration_ms=elapsed,
                ok=tool_error is None,
                summary=_summarise(name, payload),
                error=str(tool_error) if tool_error else None,
            )
        )
        return payload

    # -------------------------------------------------------------- typed api ---
    async def get_price_history(self, ticker: str, days: int) -> PriceHistory:
        return PriceHistory.model_validate(
            await self.call("get_price_history", {"ticker": ticker, "days": days})
        )

    async def get_fundamentals(self, ticker: str) -> Fundamentals:
        return Fundamentals.model_validate(await self.call("get_fundamentals", {"ticker": ticker}))

    async def compute_metrics(
        self, ticker: str, bars: list[PriceBar], pe_ratio: float | None
    ) -> Metrics:
        payload = await self.call(
            "compute_metrics",
            {
                "ticker": ticker,
                "bars": [bar.model_dump() for bar in bars],
                "pe_ratio": pe_ratio,
            },
        )
        return Metrics.model_validate(payload)

    async def fetch_rss_news(self, ticker: str, limit: int) -> NewsFeed:
        return NewsFeed.model_validate(
            await self.call("fetch_rss_news", {"ticker": ticker, "limit": limit})
        )

    # -------------------------------------------------------------- internal ---
    def _require_client(self) -> Client:
        if self._client is None:
            raise McpToolError("MCP client is not connected; use `async with McpToolClient()`")
        return self._client


def _extract_payload(result: Any) -> dict[str, Any]:
    """Pull the structured dict out of an MCP CallToolResult.

    A protocol-level tool error (``isError``) is raised rather than silently
    returned, so a server-side failure can never be mistaken for empty data.
    """
    is_error = getattr(result, "is_error", None)
    if is_error is None:
        is_error = getattr(result, "isError", False)
    if is_error:
        detail = " ".join(
            str(getattr(block, "text", "")) for block in (getattr(result, "content", []) or [])
        ).strip()
        raise McpToolError(detail or "MCP tool reported an error with no detail")

    structured = getattr(result, "structured_content", None)
    if structured is None:
        structured = getattr(result, "structuredContent", None)
    if isinstance(structured, dict):
        # MCPServer wraps non-object returns under "result"; models come through flat.
        inner = structured.get("result")
        if isinstance(inner, dict):
            return inner
        return structured

    for block in getattr(result, "content", []) or []:
        text = getattr(block, "text", None)
        if not text:
            continue
        try:
            parsed = json.loads(text)
        except json.JSONDecodeError:
            continue
        if isinstance(parsed, dict):
            return parsed

    raise McpToolError("MCP tool returned no structured payload")


def _summarise(tool: str, payload: dict[str, Any]) -> str:
    """One-line human summary rendered in the live telemetry feed."""
    if payload.get("error"):
        return str(payload["error"])[:120]
    if tool == "get_price_history":
        return f"{len(payload.get('bars', []))} daily bars"
    if tool == "get_fundamentals":
        return f"{payload.get('name') or payload.get('ticker')} · P/E {payload.get('pe_ratio')}"
    if tool == "compute_metrics":
        return (
            f"close {payload.get('last_close')} · 30d {payload.get('return_30d_pct')}% "
            f"· vol {payload.get('volatility_annualised_pct')}%"
        )
    if tool == "fetch_rss_news":
        return f"{len(payload.get('items', []))} headlines"
    return "ok"