diff --git a/README.md b/README.md index 0d93c83..f660a75 100644 --- a/README.md +++ b/README.md @@ -360,6 +360,39 @@ SessionIQ works offline with a deterministic answer engine. To upgrade: Run `.\.venv\Scripts\python.exe scripts\check_local_ai.py` to see what's active and get setup hints. Models and vector indexes download to your machine's cache — they are **never committed**. +### Optional: query your library from your own agent (MCP) + +SessionIQ speaks the [Model Context Protocol](https://modelcontextprotocol.io), so an MCP client — +Claude Desktop, Cursor, or an agent of your own — can query your library directly. + +```powershell +.\.venv\Scripts\python.exe -m pip install -e ".[mcp]" +``` + +Point the client at the server. For a client that takes a JSON config: + +```json +{ + "mcpServers": { + "sessioniq": { + "command": "/path/to/sessioniq/.venv/bin/python", + "args": ["/path/to/sessioniq/scripts/run_mcp.py"] + } + } +} +``` + +On Windows the interpreter is `.venv\Scripts\python.exe`. + +It exposes the same seven tools the in-app assistant calls — `list_projects`, `search_library`, +`filter_assets`, `compute_stat`, `asset_details`, `similar_tracks` and `next_up` — and delegates to +the same code, so a question answered here and the same question asked in the app are computed +identically. That logic is what `evals/` measures. + +Two things worth knowing. The server is **read-only**: it answers questions and changes nothing. +And it reads the library at startup, so it sees the library as of when the client launched it — +restart the client to pick up files added since. + ## 📁 Project structure ``` @@ -378,6 +411,7 @@ sessioniq/ │ ├── advisor.py # reference A/B, finish-next ranking, weekly digest │ ├── transcription.py # optional local Whisper voice-memo transcription │ ├── jobs.py # in-process background jobs (progress, cancel) +│ ├── mcp_server.py # the tool layer exposed over MCP (optional extra) │ ├── plugins.py # analyzer registry │ └── project_workspace.py# projects, smart collections, health, reports, file ops ├── web/src/ # React + TypeScript dashboard diff --git a/pyproject.toml b/pyproject.toml index 5cc36eb..e49e8a6 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -26,9 +26,13 @@ dependencies = [ [project.optional-dependencies] dev = [ + "mcp>=2.2.0", "pytest>=8.2.0", "ruff>=0.5.0", ] +mcp = [ + "mcp>=2.2.0", +] vector = [ "chromadb>=0.5.0", ] diff --git a/scripts/run_mcp.py b/scripts/run_mcp.py new file mode 100644 index 0000000..914a787 --- /dev/null +++ b/scripts/run_mcp.py @@ -0,0 +1,25 @@ +from __future__ import annotations + +import sys +from pathlib import Path + +ROOT = Path(__file__).resolve().parents[1] +sys.path.insert(0, str(ROOT / "src")) + +from sessioniq.mcp_server import LibraryUnavailable, run # noqa: E402 + +if __name__ == "__main__": + # stdout carries the protocol, so anything a human needs to read goes to + # stderr or it corrupts the stream. + try: + run() + except LibraryUnavailable as exc: + print(f"sessioniq mcp: {exc}", file=sys.stderr) + raise SystemExit(1) from exc + except ModuleNotFoundError as exc: + print( + f"sessioniq mcp: {exc}\n" + 'Install the server with: pip install -e ".[mcp]"', + file=sys.stderr, + ) + raise SystemExit(1) from exc diff --git a/src/sessioniq/mcp_server.py b/src/sessioniq/mcp_server.py new file mode 100644 index 0000000..53177d3 --- /dev/null +++ b/src/sessioniq/mcp_server.py @@ -0,0 +1,236 @@ +"""SessionIQ as a Model Context Protocol server. + +Exposes the same tool layer the in-app assistant calls, so any MCP client — +Claude Desktop, Cursor, an agent of your own — can query a local SessionIQ +library. + + python scripts/run_mcp.py # speaks MCP over stdio + +Every tool delegates to :class:`~sessioniq.tools.LibraryToolbox`. A question +answered here and the same question asked inside the app are computed by +identical code, and the evaluation in ``evals/`` covers that logic. + +Retrieval runs in memory rather than through ChromaDB. The MCP server is a +read-only query surface, so it should start instantly and must not need an +embedding model; the lexical and metadata blend it uses scored 0.733 hit@1 +against 0.767 with vectors, which is not worth a model download here. +""" + +from __future__ import annotations + +import json +from pathlib import Path +from typing import TYPE_CHECKING, Any + +from sessioniq.models import ProjectAsset +from sessioniq.project_workspace import UPLOAD_ROOT, summarize_projects +from sessioniq.retrieval import InMemoryRetriever +from sessioniq.tools import LibraryToolbox + +if TYPE_CHECKING: # The SDK is an optional install; see the `mcp` extra. + from mcp.server.mcpserver import MCPServer + +LIBRARY_INDEX_NAME = "library-index.json" + +KINDS = ("audio", "midi", "note", "image", "reference", "export") +STATUSES = ("Idea", "In Progress", "Needs Work", "Ready", "Reference", "Archived") +STAT_FIELDS = ("bpm", "rms_db", "duration", "brightness", "note_count") +STAT_OPS = ("min", "max", "avg", "count") + + +class LibraryUnavailable(RuntimeError): + """Raised when there is no library for the server to answer from.""" + + +def library_index_path() -> Path: + return UPLOAD_ROOT.parent / LIBRARY_INDEX_NAME + + +def load_library(index_path: Path | None = None) -> tuple[list[ProjectAsset], dict[str, str]]: + """Read the analyzed library from disk. + + The API keeps this in memory; the MCP server is a separate process that + reads the same file at startup. It sees the library as of the moment it + launched, which is the right trade for a query surface that should never + block on a running app. + """ + path = index_path or library_index_path() + if not path.exists(): + raise LibraryUnavailable( + f"No library found at {path}. Run the app once, or " + "`python scripts/seed_demo.py`, to create one." + ) + try: + raw = json.loads(path.read_bytes()) + except (OSError, json.JSONDecodeError) as exc: + raise LibraryUnavailable(f"Could not read {path}: {exc}") from exc + + assets: list[ProjectAsset] = [] + for item in raw.get("assets", []): + try: + assets.append(ProjectAsset.model_validate(item)) + except Exception: + # One unreadable entry must not take down the whole server; the + # app skips these on load too. + continue + + statuses = { + str(key): str(value) for key, value in (raw.get("task_statuses") or {}).items() + } + return assets, statuses + + +def build_toolbox(assets: list[ProjectAsset], statuses: dict[str, str]) -> LibraryToolbox: + """Assemble the toolbox exactly as the API does, minus the vector layer.""" + summaries = summarize_projects(assets, statuses) + tasks = [task for project in summaries for task in project.tasks] + + retriever = InMemoryRetriever() + retriever.add_assets(assets) + + return LibraryToolbox( + assets, + tasks=tasks, + search=retriever.search, + summaries=summaries, + ) + + +def build_server(index_path: Path | None = None) -> MCPServer: + """Create the MCP server, loading the library up front.""" + from mcp.server.mcpserver import MCPServer # imported here: optional install + + assets, statuses = load_library(index_path) + toolbox = build_toolbox(assets, statuses) + + def call(tool: str, arguments: dict[str, Any]) -> str: + """Run a toolbox tool and return its JSON payload.""" + return toolbox.execute(tool, arguments) + + def arguments(**values: Any) -> dict[str, Any]: + """Drop unset options so the toolbox sees only what was asked for.""" + return {key: value for key, value in values.items() if value is not None} + + server = MCPServer( + "sessioniq", + instructions=( + "Query a local SessionIQ music library: analyzed audio, MIDI and session " + "notes organized into projects. Start with list_projects to learn the " + "project names, then filter_assets or compute_stat for exact questions " + "and search_library for open-ended ones." + ), + ) + + @server.tool() + def list_projects() -> str: + """List every project with its file counts and open task count. + + Call this first: other tools accept a project_name, and this is how you + find the valid ones. + """ + return call("list_projects", {}) + + @server.tool() + def search_library(query: str, project_name: str | None = None) -> str: + """Find files by meaning across the library. + + Best for open-ended questions like "tracks that still need mastering". + For exact conditions use filter_assets, and for numbers use compute_stat. + """ + return call("search_library", arguments(query=query, project_name=project_name)) + + @server.tool() + def filter_assets( + status: str | None = None, + kind: str | None = None, + tag: str | None = None, + key: str | None = None, + bpm_min: float | None = None, + bpm_max: float | None = None, + project_name: str | None = None, + ) -> str: + """List assets matching exact conditions, with the total match count. + + Arguments are combined with AND. status is one of Idea, In Progress, + Needs Work, Ready, Reference, Archived. kind is one of audio, midi, + note, image. tag matches any part of a tag label. key is a pitch class + such as C or F#. The returned list may be truncated; total_matches + always reflects every match. + """ + return call( + "filter_assets", + arguments( + status=status, + kind=kind, + tag=tag, + key=key, + bpm_min=bpm_min, + bpm_max=bpm_max, + project_name=project_name, + ), + ) + + @server.tool() + def compute_stat( + field: str, + op: str, + status: str | None = None, + kind: str | None = None, + tag: str | None = None, + key: str | None = None, + bpm_min: float | None = None, + bpm_max: float | None = None, + project_name: str | None = None, + ) -> str: + """Compute an exact aggregate over the library. + + Use this for any count, average, or superlative — "how many tracks are + in C", "which track is the loudest". min and max also return the + winning asset. field is one of bpm, rms_db, duration, brightness, + note_count; op is one of min, max, avg, count. Filters behave as in + filter_assets. + """ + return call( + "compute_stat", + arguments( + field=field, + op=op, + status=status, + kind=kind, + tag=tag, + key=key, + bpm_min=bpm_min, + bpm_max=bpm_max, + project_name=project_name, + ), + ) + + @server.tool() + def asset_details(asset_id: str) -> str: + """Full extracted metadata for one asset, including note text and tasks. + + Use an asset_id returned by another tool. + """ + return call("asset_details", {"asset_id": asset_id}) + + @server.tool() + def similar_tracks(asset_id: str, limit: int = 5) -> str: + """Tracks most similar to one asset, across tempo, key, tone and loudness.""" + return call("similar_tracks", {"asset_id": asset_id, "limit": limit}) + + @server.tool() + def next_up() -> str: + """Rank projects by what is closest to finishable right now. + + Blends readiness, task progress and freshness, and flags + near-finished-but-stale work as stalled. Use for "what should I + finish next?". + """ + return call("next_up", {}) + + return server + + +def run(index_path: Path | None = None) -> None: + """Serve over stdio until the client disconnects.""" + build_server(index_path).run(transport="stdio") diff --git a/src/sessioniq/tools.py b/src/sessioniq/tools.py index 873af60..0dd2656 100644 --- a/src/sessioniq/tools.py +++ b/src/sessioniq/tools.py @@ -212,6 +212,39 @@ def _brief(asset: ProjectAsset) -> dict: } +# Payload keys holding lists of results, in the order worth trimming. +_TRUNCIBLE_KEYS = ("assets", "matches", "ranking", "projects") + + +def _shrink_payload(result: object, limit: int) -> object: + """Trim a tool payload until it serialises under ``limit``. + + Chopping the serialised string produced invalid JSON, which the model had + to cope with and an MCP client would simply fail to parse. Dropping whole + list items keeps the payload valid and still bounds its size. + """ + if not isinstance(result, dict): + return result + + candidate = result + while len(json.dumps(candidate, default=str)) > limit: + for key in _TRUNCIBLE_KEYS: + items = candidate.get(key) + if isinstance(items, list) and len(items) > 1: + candidate = {**candidate, key: items[: len(items) // 2], "truncated": True} + break + else: + shortened = { + key: (value[:400] if isinstance(value, str) and len(value) > 400 else value) + for key, value in candidate.items() + } + if shortened != candidate: + candidate = {**shortened, "truncated": True} + continue + return {"truncated": True, "error": "Result too large to return in full."} + return candidate + + class LibraryToolbox: """Executes tool calls against a snapshot of the library.""" @@ -254,10 +287,7 @@ def execute(self, name: str, arguments: dict) -> str: result = handler(arguments) except Exception as exc: # One bad argument must not kill the answer. result = {"error": f"{type(exc).__name__}: {exc}"} - text = json.dumps(result, default=str) - if len(text) > MAX_TOOL_RESULT_CHARS: - text = text[:MAX_TOOL_RESULT_CHARS] + '…"}' - return text + return json.dumps(_shrink_payload(result, MAX_TOOL_RESULT_CHARS), default=str) # --- Filter shared by filter_assets and compute_stat ----------------- diff --git a/tests/test_mcp_server.py b/tests/test_mcp_server.py new file mode 100644 index 0000000..7586ea6 --- /dev/null +++ b/tests/test_mcp_server.py @@ -0,0 +1,176 @@ +"""Tests for the MCP server. + +The SDK is an optional extra, so the server-surface tests skip when it is not +installed. Library loading and toolbox wiring are tested regardless. +""" + +from __future__ import annotations + +import asyncio +import importlib.util +import json +from pathlib import Path + +import pytest + +from sessioniq.mcp_server import ( + LibraryUnavailable, + build_server, + build_toolbox, + load_library, +) + +requires_mcp = pytest.mark.skipif( + importlib.util.find_spec("mcp") is None, + reason="the mcp extra is not installed", +) + +EXPECTED_TOOLS = { + "list_projects", + "search_library", + "filter_assets", + "compute_stat", + "asset_details", + "similar_tracks", + "next_up", +} + + +def write_library(directory: Path, payload: dict | None = None) -> Path: + directory.mkdir(parents=True, exist_ok=True) + path = directory / "library-index.json" + body = json.dumps(payload if payload is not None else SAMPLE_LIBRARY) + path.write_text(body, encoding="utf-8") + return path + + +SAMPLE_LIBRARY = { + "assets": [ + { + "id": "loud", + "file_name": "loud.wav", + "kind": "audio", + "project_name": "Album/One", + "status": "Ready", + "tags": [{"label": "single", "ai_suggested": False}], + "audio": { + "bpm_estimate": 120.0, + "key_estimate": "C", + "rms_db": -9.0, + "spectral_centroid_mean": 1200.0, + "duration_seconds": 180.0, + }, + }, + { + "id": "quiet", + "file_name": "quiet.wav", + "kind": "audio", + "project_name": "Album/Two", + "status": "Idea", + "tags": [], + "audio": { + "bpm_estimate": 90.0, + "key_estimate": "A", + "rms_db": -20.0, + "spectral_centroid_mean": 800.0, + "duration_seconds": 200.0, + }, + }, + ], + "task_statuses": {}, +} + + +def payload_of(result) -> dict: + return json.loads(result.content[0].text) + + +class TestLibraryLoading: + def test_missing_library_is_reported_clearly(self, tmp_path): + with pytest.raises(LibraryUnavailable) as excinfo: + load_library(tmp_path / "nope.json") + assert "seed_demo" in str(excinfo.value) + + def test_reads_assets_and_statuses(self, tmp_path): + path = write_library(tmp_path, {**SAMPLE_LIBRARY, "task_statuses": {"t": "Done"}}) + assets, statuses = load_library(path) + assert [asset.id for asset in assets] == ["loud", "quiet"] + assert statuses == {"t": "Done"} + + def test_unreadable_entries_are_skipped_not_fatal(self, tmp_path): + payload = { + "assets": [SAMPLE_LIBRARY["assets"][0], {"id": "broken"}], + "task_statuses": {}, + } + assets, _ = load_library(write_library(tmp_path, payload)) + assert [asset.id for asset in assets] == ["loud"] + + def test_corrupt_json_is_reported_clearly(self, tmp_path): + tmp_path.mkdir(parents=True, exist_ok=True) + path = tmp_path / "library-index.json" + path.write_text("{not json", encoding="utf-8") + with pytest.raises(LibraryUnavailable): + load_library(path) + + +class TestToolboxWiring: + def test_search_is_wired_to_the_retriever(self, tmp_path): + assets, statuses = load_library(write_library(tmp_path)) + toolbox = build_toolbox(assets, statuses) + result = json.loads(toolbox.execute("search_library", {"query": "loud"})) + assert [match["file_name"] for match in result["matches"]][:1] == ["loud.wav"] + + def test_summaries_are_available_to_next_up(self, tmp_path): + assets, statuses = load_library(write_library(tmp_path)) + toolbox = build_toolbox(assets, statuses) + result = json.loads(toolbox.execute("next_up", {})) + assert "error" not in result, "next_up needs project summaries to be wired" + + +@requires_mcp +class TestServerSurface: + def test_exposes_the_tool_layer(self, tmp_path): + server = build_server(write_library(tmp_path)) + tools = asyncio.run(server.list_tools()) + assert {tool.name for tool in tools} == EXPECTED_TOOLS + + def test_every_tool_describes_itself(self, tmp_path): + server = build_server(write_library(tmp_path)) + for tool in asyncio.run(server.list_tools()): + assert tool.description, f"{tool.name} has no description for the model to read" + + +@requires_mcp +class TestToolDelegation: + """The MCP layer must not reimplement anything the app already decides.""" + + def test_compute_stat_returns_the_winning_asset(self, tmp_path): + server = build_server(write_library(tmp_path)) + result = asyncio.run( + server.call_tool("compute_stat", {"field": "rms_db", "op": "max"}) + ) + assert result.is_error is False + assert payload_of(result)["asset"]["file_name"] == "loud.wav" + + def test_filter_assets_honours_conditions(self, tmp_path): + server = build_server(write_library(tmp_path)) + result = asyncio.run(server.call_tool("filter_assets", {"status": "Ready"})) + body = payload_of(result) + assert body["total_matches"] == 1 + assert body["assets"][0]["file_name"] == "loud.wav" + + def test_list_projects_counts_files(self, tmp_path): + server = build_server(write_library(tmp_path)) + result = asyncio.run(server.call_tool("list_projects", {})) + projects = {item["project"]: item["files"] for item in payload_of(result)["projects"]} + assert projects == {"Album/One": 1, "Album/Two": 1} + + def test_unknown_asset_id_is_an_error_not_a_crash(self, tmp_path): + server = build_server(write_library(tmp_path)) + result = asyncio.run(server.call_tool("asset_details", {"asset_id": "nope"})) + assert "error" in payload_of(result) + + def test_read_only_tools_need_no_arguments(self, tmp_path): + server = build_server(write_library(tmp_path)) + for name in ("list_projects", "next_up"): + assert payload_of(asyncio.run(server.call_tool(name, {}))) diff --git a/tests/test_tools.py b/tests/test_tools.py index 4e979db..e0ee3fb 100644 --- a/tests/test_tools.py +++ b/tests/test_tools.py @@ -11,7 +11,7 @@ ProjectTask, TaskStatus, ) -from sessioniq.tools import LibraryToolbox +from sessioniq.tools import MAX_TOOL_RESULT_CHARS, LibraryToolbox, _shrink_payload def _audio( @@ -173,3 +173,40 @@ def test_unknown_tool_and_bad_arguments_return_errors(): assert "error" in _run(toolbox, "does_not_exist", {}) # Non-serializable filter values must not raise out of execute(). assert "error" in _run(toolbox, "compute_stat", {"field": "bpm", "op": "explode"}) + + +class TestPayloadShrinking: + """Over-long results must stay parseable. + + The previous implementation chopped the serialised string and appended a + marker, which produced invalid JSON. The model coped with it; an MCP client + cannot, so the payload is now trimmed structurally instead. + """ + + def test_small_payloads_pass_through_untouched(self): + payload = {"assets": [{"file_name": "a.wav"}]} + assert _shrink_payload(payload, MAX_TOOL_RESULT_CHARS) == payload + + def test_a_long_list_is_halved_and_still_parses(self): + payload = {"assets": [{"file_name": f"{i}.wav", "note": "x" * 80} for i in range(60)]} + shrunk = _shrink_payload(payload, MAX_TOOL_RESULT_CHARS) + assert shrunk["truncated"] is True + assert 0 < len(shrunk["assets"]) < 60 + assert len(json.dumps(shrunk)) <= MAX_TOOL_RESULT_CHARS + + def test_oversized_strings_shrink_when_no_list_can(self): + shrunk = _shrink_payload({"metadata": "y" * 9000}, MAX_TOOL_RESULT_CHARS) + assert shrunk["truncated"] is True + assert len(json.dumps(shrunk)) <= MAX_TOOL_RESULT_CHARS + + def test_an_impossible_payload_reports_instead_of_emitting_broken_json(self): + shrunk = _shrink_payload({"assets": [{"k": "z" * 9000}]}, 200) + assert shrunk["truncated"] is True + json.loads(json.dumps(shrunk)) + + def test_execute_stays_parseable_over_a_large_library(self): + # filter_assets caps its own list, so a large result is exactly where + # the old string-chopping corruption used to surface. + toolbox = _toolbox([_audio(f"a{i}", 100.0 + i) for i in range(60)]) + result = _run(toolbox, "filter_assets", {}) + assert result["total_matches"] == 60