diff --git a/lib/anubis/mcp/error.ex b/lib/anubis/mcp/error.ex index c0bb1507..698ab70c 100644 --- a/lib/anubis/mcp/error.ex +++ b/lib/anubis/mcp/error.ex @@ -259,7 +259,7 @@ defmodule Anubis.MCP.Error do } |> Enum.reject(fn {_, v} -> is_nil(v) end) |> Map.new() - |> then(&%{"error" => &1, "id" => id}) + |> then(&%{"jsonrpc" => "2.0", "error" => &1, "id" => id}) end # Private helpers diff --git a/lib/anubis/server.ex b/lib/anubis/server.ex index b1088515..fb1ee98d 100644 --- a/lib/anubis/server.ex +++ b/lib/anubis/server.ex @@ -86,6 +86,18 @@ defmodule Anubis.Server do Most protocol handling is automatic - you typically only implement `init/2` for setup and occasionally override other callbacks for custom behavior. + + ## Sending Notifications + + Notification functions use `send(self(), ...)` and must be called from within the + Session process (i.e., inside callbacks). For sending from external processes or tasks, + use `send/2` with the session PID directly. + + # Inside a callback: + def handle_info(:data_changed, frame) do + Anubis.Server.send_tools_list_changed() + {:noreply, frame} + end """ alias Anubis.Server.Component @@ -112,15 +124,12 @@ defmodule Anubis.Server do This callback is invoked while the MCP handshake starts and so the client may not sent the `notifications/initialized` message yet. For checking if the notification was already sent - and the MCP handshake was successfully completed, you can call the `initialized?/1` function. + and the MCP handshake was successfully completed, you can check the `context.initialized` field + in the frame. It receives the client's information and the current frame, allowing you to perform client-specific setup, validate capabilities, or prepare resources based on the connected client. - - The client_info parameter contains details about the connected client including its - name, version, and any additional metadata. Use this to tailor your server's behavior - to specific client implementations or versions. """ @callback init(client_info :: map(), Frame.t()) :: {:ok, Frame.t()} @@ -128,12 +137,7 @@ defmodule Anubis.Server do Handles a tool call request. This callback is invoked when a client calls a specific tool. It receives the tool name, - the arguments provided by the client, and the current frame. Developers's implementation should - execute the tool's logic and return the result. - - This callback handles both module-based components (registered with `component`) and - runtime components (registered with `Frame.register_tool/3`). For module-based tools, - the framework automatically generates pattern-matched clauses during compilation. + the arguments provided by the client, and the current frame. """ @callback handle_tool_call(name :: String.t(), arguments :: map(), Frame.t()) :: {:reply, result :: term(), Frame.t()} @@ -141,14 +145,6 @@ defmodule Anubis.Server do @doc """ Handles a resource read request. - - This callback is invoked when a client requests to read a specific resource. It receives - the resource URI and the current frame. Developer's implementation should retrieve and return - the resource content. - - This callback handles both module-based components (registered with `component`) and - runtime components (registered with `Frame.register_resource/3`). For module-based resources, - the framework automatically generates pattern-matched clauses during compilation. """ @callback handle_resource_read(uri :: String.t(), Frame.t()) :: {:reply, content :: map(), Frame.t()} @@ -156,13 +152,6 @@ defmodule Anubis.Server do @doc """ Handles a prompt get request. - - This callback is invoked when a client requests a specific prompt template. It receives - the prompt name, any arguments to fill into the template, and the current frame. - - This callback handles both module-based components (registered with `component`) and - runtime components (registered with `Frame.register_prompt/3`). For module-based prompts, - the framework automatically generates pattern-matched clauses during compilation. """ @callback handle_prompt_get(name :: String.t(), arguments :: map(), Frame.t()) :: {:reply, messages :: list(), Frame.t()} @@ -171,19 +160,7 @@ defmodule Anubis.Server do @doc """ Low-level handler for any MCP request. - This is an advanced callback that gives you complete control over request handling. - When implemented, it bypasses the automatic routing to `handle_tool_call/3`, - `handle_resource_read/2`, and `handle_prompt_get/3` and all other requests that are - handled internally, like `tools/list` and `logging/setLevel`. - - Use this when you need to: - - Implement custom request methods beyond the standard MCP protocol - - Add middleware-like processing before requests reach specific handlers - - Override the framework's default request routing behavior - - Note: If you implement this callback, you become responsible for handling ALL - MCP requests, including standard protocol methods like `tools/list`, `resources/list`, etc. - Consider using the specific callbacks instead unless you need this level of control. + When implemented, it bypasses automatic routing to specific handlers. """ @callback handle_request(request :: request(), state :: Frame.t()) :: {:reply, response :: response(), new_state :: Frame.t()} @@ -192,99 +169,20 @@ defmodule Anubis.Server do @doc """ Handles incoming MCP notifications from clients. - - Notifications are one-way messages in the MCP protocol - the client informs the server - about events or state changes without expecting a response. This fire-and-forget pattern - is perfect for status updates, progress tracking, and lifecycle events. - - **Standard MCP Notifications from Clients:** - - `notifications/initialized` - Client signals it's ready after successful initialization - - `notifications/cancelled` - Client requests cancellation of an in-progress operation - - `notifications/progress` - Client reports progress on a long-running operation - - `notifications/roots/list_changed` - Client's available filesystem roots have changed - - Unlike requests, notifications never receive responses. Any errors during processing - are typically logged but not communicated back to the client. This makes notifications - ideal for optional features like progress tracking where delivery isn't guaranteed. - - The server processes these notifications to update its internal state, trigger side effects, - or coordinate with other parts of the system. When using `use Anubis.Server`, basic - notification handling is provided, but you'll often want to override this callback - to handle progress updates or cancellations specific to your server's operations. """ @callback handle_notification(notification :: notification(), state :: Frame.t()) :: {:noreply, new_state :: Frame.t()} | {:error, error :: mcp_error(), new_state :: Frame.t()} - @doc """ - Provides the server's identity information during initialization. - - This callback is called during the MCP handshake to identify your server to connecting clients. - The information returned here helps clients understand which server they're talking to and - ensures version compatibility. - - When using `use Anubis.Server`, this callback is automatically implemented using the - `name` and `version` options you provide. You only need to implement this manually if - you require dynamic server information based on runtime conditions. - """ @callback server_info :: server_info() - - @doc """ - Declares the server's capabilities during initialization. - - This callback tells clients what features your server supports - which types of resources - it can provide, what tools it can execute, whether it supports logging configuration, etc. - The capabilities you declare here directly impact which requests the client will send. - - When using `use Anubis.Server` with the `capabilities` option, this callback is automatically - implemented based on your configuration. The macro analyzes your registered components and - builds the appropriate capability map, so you rarely need to implement this manually. - """ @callback server_capabilities :: server_capabilities() - - @doc """ - Specifies which MCP protocol versions this server can speak. - - Protocol version negotiation ensures client and server can communicate effectively. - During initialization, the client and server agree on a mutually supported version. - This callback returns the list of versions your server understands, typically in - order of preference from newest to oldest. - - When using `use Anubis.Server`, this is automatically implemented with sensible defaults - covering current and recent protocol versions. Override only if you need to restrict - or extend version support for specific compatibility requirements. - """ @callback supported_protocol_versions() :: [String.t()] - @doc """ - Handles non-MCP messages sent to the server process. - - While `handle_request` and `handle_notification` deal with MCP protocol messages, - this callback handles everything else - timer events, messages from other processes, - system signals, and any custom inter-process communication your server needs. - - This is particularly useful for servers that need to react to external events - (like file system changes or database updates) and notify connected clients through - MCP notifications. Think of it as the bridge between your Elixir application's - internal events and the MCP protocol's notification system. - """ @callback handle_info(event :: term, Frame.t()) :: {:noreply, Frame.t()} | {:noreply, Frame.t(), timeout() | :hibernate | {:continue, arg :: term}} | {:stop, reason :: term, Frame.t()} - @doc """ - Handles synchronous calls to the server process. - - This optional callback allows you to handle custom synchronous calls made to your - MCP server process using `GenServer.call/2`. This is useful for implementing - administrative functions, status queries, or any synchronous operations that - need to interact with the server's internal state. - - The callback follows standard GenServer semantics and should return appropriate - reply tuples. If not implemented, the Base module provides a default implementation - that handles standard MCP operations. - """ @callback handle_call(request :: term, from :: GenServer.from(), Frame.t()) :: {:reply, reply :: term, Frame.t()} | {:reply, reply :: term, Frame.t(), timeout() | :hibernate | {:continue, arg :: term}} @@ -293,64 +191,13 @@ defmodule Anubis.Server do | {:stop, reason :: term, reply :: term, Frame.t()} | {:stop, reason :: term, Frame.t()} - @doc """ - Handles asynchronous casts to the server process. - - This optional callback allows you to handle custom asynchronous messages sent to your - MCP server process using `GenServer.cast/2`. This is useful for fire-and-forget - operations, background tasks, or any asynchronous operations that don't require - an immediate response. - - The callback follows standard GenServer semantics. If not implemented, the Base - module provides a default implementation that handles standard MCP operations. - """ @callback handle_cast(request :: term, Frame.t()) :: {:noreply, Frame.t()} | {:noreply, Frame.t(), timeout() | :hibernate | {:continue, arg :: term}} | {:stop, reason :: term, Frame.t()} - @doc """ - Cleans up when the server process terminates. - - This optional callback is invoked when the server process is about to terminate. - It allows you to perform cleanup operations, close connections, save state, - or release resources before the process exits. - - The callback receives the termination reason and the current frame. Any return - value is ignored. If not implemented, the Base module provides a default - implementation that logs the termination event. - """ @callback terminate(reason :: term, Frame.t()) :: term - @doc """ - Handles the response from a sampling/createMessage request sent to the client. - - This callback is invoked when the client responds to a sampling request initiated - by the server. The response contains the generated message from the client's LLM. - - ## Parameters - - * `response` - The response from the client containing: - * `"role"` - The role of the generated message (typically "assistant") - * `"content"` - The content object with type and data - * `"model"` - The model used for generation - * `"stopReason"` - Why generation stopped (e.g., "endTurn") - * `request_id` - The ID of the original request for correlation - * `frame` - The current server frame - - ## Returns - - * `{:noreply, frame}` - Continue processing - * `{:stop, reason, frame}` - Stop the server - - ## Examples - - def handle_sampling(response, request_id, frame) do - %{"content" => %{"text" => text}} = response - # Process the generated text... - {:noreply, frame} - end - """ @callback handle_sampling( response :: map(), request_id :: String.t(), @@ -359,25 +206,10 @@ defmodule Anubis.Server do {:noreply, Frame.t()} | {:stop, reason :: term(), Frame.t()} - @doc """ - Handles completion requests from the client. - - This callback is invoked when a client requests completions for a reference. - The reference indicates what type of completion is being requested. - - Note: This callback will only be invoked if user declared the `completion` capability - on server definition - """ @callback handle_completion(ref :: String.t(), argument :: map(), Frame.t()) :: {:reply, Response.t() | map(), Frame.t()} | {:error, mcp_error(), Frame.t()} - @doc """ - Handles the response from a roots/list request sent to the client. - - This callback is invoked when the client responds to a roots list request - initiated by the server. The response contains the available root URIs. - """ @callback handle_roots( roots :: list(map()), request_id :: String.t(), @@ -400,28 +232,6 @@ defmodule Anubis.Server do handle_completion: 3, handle_roots: 3 - @doc """ - Checks if the MCP session has been initialized. - - Returns true if the client has completed the initialization handshake and sent - the `notifications/initialized` message. This is useful for guarding operations - that require an active session. - - ## Examples - - def handle_info(:check_status, frame) do - if Anubis.Server.initialized?(frame) do - # Perform operations requiring initialized session - {:noreply, frame} - else - # Wait for initialization - {:noreply, frame} - end - end - """ - @spec initialized?(Frame.t()) :: boolean() - def initialized?(%Frame{initialized: initialized}), do: initialized - @doc false defguard is_server_capability(capability) when capability in @server_capabilities @@ -462,14 +272,6 @@ defmodule Anubis.Server do @doc """ Registers a component (tool, prompt, or resource) with the server. - - ## Examples - - # Register with auto-derived name - component MyServer.Tools.Calculator - - # Register with custom name - component MyServer.Tools.FileManager, name: "files" """ defmacro component(module, opts \\ []) do quote bind_quoted: [module: module, opts: opts] do @@ -734,75 +536,68 @@ defmodule Anubis.Server do def validate_server_info!(_, name, version) when is_binary(name) and is_binary(version), do: :ok - # Notification Functions + # Notification Functions — all use send(self(), ...) to the current Session process @doc """ - Sends a resources list changed notification to connected clients. + Sends a resources list changed notification. - Use this when the available resources have changed (added, removed, or modified). - The client will typically re-fetch the resource list in response. + **Must be called from within a Session callback** — the current process must be + the Session GenServer. Calling from outside a callback will silently lose the message. + + For external processes, use `send(session_pid, {:send_notification, "notifications/resources/list_changed", %{}})`. """ - @spec send_resources_list_changed(Frame.t()) :: :ok - def send_resources_list_changed(%Frame{} = frame) do - queue_notification(frame, "notifications/resources/list_changed", %{}) + @spec send_resources_list_changed :: :ok + def send_resources_list_changed do + send(self(), {:send_notification, "notifications/resources/list_changed", %{}}) + :ok end @doc """ Sends a resource updated notification for a specific resource. - Use this when the content of a specific resource has changed. - Clients that have subscribed to this resource will be notified. + **Must be called from within a Session callback** — see `send_resources_list_changed/0` for details. """ - @spec send_resource_updated( - Frame.t(), - uri :: String.t(), - timestamp :: DateTime.t() | nil - ) :: - :ok - def send_resource_updated(%Frame{} = frame, uri, timestamp \\ nil) do + @spec send_resource_updated(uri :: String.t(), timestamp :: DateTime.t() | nil) :: :ok + def send_resource_updated(uri, timestamp \\ nil) do params = %{"uri" => uri} params = if timestamp, do: Map.put(params, "timestamp", timestamp), else: params - queue_notification(frame, "notifications/resources/updated", params) + send(self(), {:send_notification, "notifications/resources/updated", params}) + :ok end @doc """ - Sends a prompts list changed notification to connected clients. + Sends a prompts list changed notification. - Use this when the available prompts have changed (added, removed, or modified). - The client will typically re-fetch the prompt list in response. + **Must be called from within a Session callback** — see `send_resources_list_changed/0` for details. """ - @spec send_prompts_list_changed(Frame.t()) :: :ok - def send_prompts_list_changed(%Frame{} = frame) do - queue_notification(frame, "notifications/prompts/list_changed", %{}) + @spec send_prompts_list_changed :: :ok + def send_prompts_list_changed do + send(self(), {:send_notification, "notifications/prompts/list_changed", %{}}) + :ok end @doc """ - Sends a tools list changed notification to connected clients. + Sends a tools list changed notification. - Use this when the available tools have changed (added, removed, or modified). - The client will typically re-fetch the tool list in response. + **Must be called from within a Session callback** — see `send_resources_list_changed/0` for details. """ - @spec send_tools_list_changed(Frame.t()) :: :ok - def send_tools_list_changed(%Frame{} = frame) do - queue_notification(frame, "notifications/tools/list_changed", %{}) + @spec send_tools_list_changed :: :ok + def send_tools_list_changed do + send(self(), {:send_notification, "notifications/tools/list_changed", %{}}) + :ok end @doc """ Sends a log message to the client. - Use this to send diagnostic or informational messages to the client's logging system. + **Must be called from within a Session callback** — see `send_resources_list_changed/0` for details. """ - @spec send_log_message( - Frame.t(), - level :: Logger.level(), - message :: String.t(), - metadata :: map() | nil - ) :: :ok - def send_log_message(%Frame{} = frame, level, message, data \\ nil) do + @spec send_log_message(level :: Logger.level(), message :: String.t(), metadata :: map() | nil) :: :ok + def send_log_message(level, message, data \\ nil) do params = %{"level" => level, "message" => message} params = if data, do: Map.put(params, "data", data), else: params - - queue_notification(frame, "notifications/log/message", params) + send(self(), {:send_notification, "notifications/log/message", params}) + :ok end @type progress_token :: String.t() | non_neg_integer @@ -811,66 +606,34 @@ defmodule Anubis.Server do @doc """ Sends a progress notification for an ongoing operation. - - Use this to update the client on the progress of long-running operations. """ - @spec send_progress(Frame.t(), progress_token, progress_step, opts) :: :ok + @spec send_progress(progress_token, progress_step, opts) :: :ok when opts: list({:total, progress_total} | {:message, String.t()}) - def send_progress(%Frame{} = frame, progress_token, progress, opts \\ []) do + def send_progress(progress_token, progress, opts \\ []) do total = opts[:total] message = opts[:message] params = %{"progressToken" => progress_token, "progress" => progress} params = if total, do: Map.put(params, "total", total), else: params params = if message, do: Map.put(params, "message", message), else: params - - queue_notification(frame, "notifications/progress", params) - end - - defp queue_notification(frame, method, params) do - registry = frame.private.server_registry - server = frame.private.server_module - pid = registry.whereis_server(server) - send(pid, {:send_notification, method, params}) + send(self(), {:send_notification, "notifications/progress", params}) :ok end - # Sampling Request Functions - @doc """ Sends a sampling/createMessage request to the client. - This function is used when the server needs the client to generate a message - using its language model. The client must have declared the sampling capability - during initialization. - - Note: This is an asynchronous operation. The response will be delivered to your + This is an asynchronous operation. The response will be delivered to your `handle_sampling/3` callback. - - Check https://modelcontextprotocol.io/specification/2025-06-18/client/sampling for more information - - ## Examples - - messages = [ - %{"role" => "user", "content" => %{"type" => "text", "text" => "Hello"}} - ] - - model_preferences = %{"costPriority" => 1.0, "speedPriority" => 0.1, "hints" => [%{"name" => "claude"}]} - - :ok = Anubis.Server.send_sampling_request(frame, messages, - model_preferences: model_preferences, - system_prompt: "You are a helpful assistant", - max_tokens: 100 - ) """ - @spec send_sampling_request(Frame.t(), list(map()), configuration) :: :ok + @spec send_sampling_request(list(map()), configuration) :: :ok when configuration: list( - {:model_preferences, map | nil} + {:model_preferences, map() | nil} | {:system_prompt, String.t() | nil} - | {:max_token, non_neg_integer | nil} - | {:timeout, non_neg_integer | nil} + | {:max_tokens, non_neg_integer() | nil} + | {:timeout, non_neg_integer() | nil} ) - def send_sampling_request(%Frame{} = frame, messages, opts \\ []) when is_list(messages) do + def send_sampling_request(messages, opts \\ []) when is_list(messages) do params = %{"messages" => messages} params = @@ -883,27 +646,17 @@ defmodule Anubis.Server do end) timeout = Keyword.get(opts, :timeout, 30_000) - registry = frame.private.server_registry - server = frame.private.server_module - pid = registry.whereis_server(server) - send(pid, {:send_sampling_request, params, timeout}) + send(self(), {:send_sampling_request, params, timeout}) :ok end @doc """ Sends a roots/list request to the client. - - This function queries the client for available root URIs. The client must have - declared the roots capability during initialization. """ - @spec send_roots_request(Frame.t(), list({:timeout, non_neg_integer | nil})) :: :ok - def send_roots_request(%Frame{} = frame, opts \\ []) do + @spec send_roots_request(list({:timeout, non_neg_integer() | nil})) :: :ok + def send_roots_request(opts \\ []) do timeout = Keyword.get(opts, :timeout, 30_000) - - registry = frame.private.server_registry - server = frame.private.server_module - pid = registry.whereis_server(server) - send(pid, {:send_roots_request, timeout}) + send(self(), {:send_roots_request, timeout}) :ok end end diff --git a/lib/anubis/server/base.ex b/lib/anubis/server/base.ex deleted file mode 100644 index f373bad5..00000000 --- a/lib/anubis/server/base.ex +++ /dev/null @@ -1,910 +0,0 @@ -defmodule Anubis.Server.Base do - @moduledoc false - - use GenServer - use Anubis.Logging - - import Peri - - alias Anubis.MCP.Error - alias Anubis.MCP.ID - alias Anubis.MCP.Message - alias Anubis.Server - alias Anubis.Server.Frame - alias Anubis.Server.Session - alias Anubis.Server.Session.Supervisor, as: SessionSupervisor - alias Anubis.Telemetry - - require Message - require Server - require Session - - @default_session_idle_timeout to_timeout(minute: 30) - - @type t :: %{ - module: module, - server_info: map, - capabilities: map, - frame: Frame.t(), - supported_versions: list(String.t()), - transport: [layer: module, name: GenServer.name()], - registry: module, - sessions: %{required(String.t()) => {GenServer.name(), reference()}}, - session_idle_timeout: pos_integer(), - expiry_timers: %{required(String.t()) => reference()}, - server_requests: %{ - required(String.t()) => %{ - method: String.t(), - session_id: String.t(), - metadata: map(), - timer_ref: reference() - } - } - } - - @typedoc """ - MCP server options - - - `:module` - The module implementing the server behavior (required) - - `:name` - Optional name for registering the GenServer - - `:session_idle_timeout` - Time in milliseconds before idle sessions expire (default: 30 minutes) - """ - @type option :: - {:module, GenServer.name()} - | {:name, GenServer.name()} - | {:session_idle_timeout, pos_integer()} - | GenServer.option() - - defschema(:parse_options, [ - {:module, {:required, {:custom, &Anubis.genserver_name/1}}}, - {:name, {:required, {:custom, &Anubis.genserver_name/1}}}, - {:transport, {:required, {:custom, &Anubis.server_transport/1}}}, - {:registry, {:atom, {:default, Anubis.Server.Registry}}}, - {:session_idle_timeout, {{:integer, {:gte, 1}}, {:default, @default_session_idle_timeout}}}, - {:timeout, {:integer, {:default, to_timeout(second: 30)}}} - ]) - - @spec start_link(Enumerable.t(option())) :: GenServer.on_start() - def start_link(opts) do - opts = parse_options!(opts) - server_name = Keyword.fetch!(opts, :name) - - GenServer.start_link(__MODULE__, Map.new(opts), name: server_name) - end - - # GenServer callbacks - - @impl GenServer - def init(%{module: module} = opts) do - server_info = module.server_info() - capabilities = module.server_capabilities() - protocol_versions = module.supported_protocol_versions() - - state = %{ - module: module, - server_info: server_info, - capabilities: capabilities, - supported_versions: protocol_versions, - transport: Map.new(opts.transport), - registry: opts.registry, - sessions: %{}, - session_idle_timeout: opts.session_idle_timeout, - expiry_timers: %{}, - frame: Frame.new(), - server_requests: %{}, - timeout: opts.timeout - } - - Logging.server_event("starting", %{ - module: module, - server_info: server_info, - capabilities: capabilities - }) - - Telemetry.execute( - Telemetry.event_server_init(), - %{system_time: System.system_time()}, - %{module: module, server_info: server_info, capabilities: capabilities} - ) - - {:ok, state, :hibernate} - end - - @impl GenServer - def handle_call({:request, decoded, session_id, context}, _from, state) when is_map(decoded) do - with {:ok, {%Session{} = session, state}} <- - maybe_attach_session(session_id, context, state) do - case handle_single_request(decoded, session, state) do - {:reply, {:ok, %{"result" => result} = response}, new_state} -> - request_id = response["id"] - if request_id, do: Session.complete_request(session.name, request_id) - - {:reply, Message.encode_response(%{"result" => result}, response["id"]), new_state} - - {:reply, {:ok, %{"error" => error} = response}, new_state} -> - request_id = response["id"] - if request_id, do: Session.complete_request(session.name, request_id) - - {:reply, Message.encode_error(%{"error" => error}, response["id"]), new_state} - - {:reply, {:error, error}, new_state} -> - request_id = decoded["id"] - if request_id, do: Session.complete_request(session.name, request_id) - {:reply, {:error, error}, new_state} - end - end - end - - def handle_call(request, from, %{module: module} = state) do - case module.handle_call(request, from, state.frame) do - {:reply, reply, frame} -> - {:reply, reply, %{state | frame: frame}} - - {:reply, reply, frame, cont} -> - {:reply, reply, %{state | frame: frame}, cont} - - {:noreply, frame} -> - {:noreply, %{state | frame: frame}} - - {:noreply, frame, cont} -> - {:noreply, %{state | frame: frame}, cont} - - {:stop, reason, reply, frame} -> - {:stop, reason, reply, %{state | frame: frame}} - - {:stop, reason, frame} -> - {:stop, reason, %{state | frame: frame}} - end - end - - @impl GenServer - def handle_cast({:notification, decoded, session_id, context}, state) when is_map(decoded) do - with {:ok, {%Session{} = session, state}} <- - maybe_attach_session(session_id, context, state) do - if Message.is_initialize_lifecycle(decoded) or Session.is_initialized(session) do - handle_notification(decoded, session, state) - else - Logging.server_event("session_not_initialized_check", %{ - session_id: session.id, - initialized: session.initialized, - method: decoded["method"] - }) - - {:noreply, state} - end - end - end - - def handle_cast({:response, decoded, _session_id, _context}, state) when is_map(decoded) do - cond do - Message.is_response(decoded) and server_request?(decoded["id"], state) -> - handle_server_request_response(decoded, state) - - Message.is_error(decoded) and server_request?(decoded["id"], state) -> - handle_server_request_error(decoded, state) - - true -> - Logging.server_event( - "unexpected_response", - %{message: decoded}, - level: :warning - ) - - {:noreply, state} - end - end - - def handle_cast(request, %{module: module} = state) do - case module.handle_cast(request, state.frame) do - {:noreply, frame} -> {:noreply, %{state | frame: frame}} - {:noreply, frame, cont} -> {:noreply, %{state | frame: frame}, cont} - {:stop, reason, frame} -> {:stop, reason, %{state | frame: frame}} - end - end - - @impl GenServer - def handle_info({:DOWN, ref, :process, _pid, reason}, state) do - session_entry = - Enum.find(state.sessions, fn - {_id, {_name, ^ref}} -> true - _ -> false - end) - - case session_entry do - {session_id, _} -> - Logging.server_event("session_terminated", %{ - session_id: session_id, - reason: reason - }) - - sessions = Map.delete(state.sessions, session_id) - state = cancel_session_expiry(session_id, state) - frame = state.frame - - frame = - if frame.private[:session_id] == session_id, - do: Frame.clear_session(frame), - else: frame - - {:noreply, %{state | sessions: sessions, frame: frame}} - - nil -> - {:noreply, state} - end - end - - def handle_info({:send_notification, method, params}, state) do - with {:ok, notification} <- encode_notification(method, params), - :ok <- send_to_transport(state.transport, notification, timeout: state.timeout) do - {:noreply, state} - else - {:error, err} -> - Logging.server_event("failed_send_notification", %{method: method, error: err}, level: :error) - - {:noreply, state} - end - end - - def handle_info({:session_expired, session_id}, state) do - if Map.get(state.sessions, session_id) do - Logging.server_event("session_expired", %{session_id: session_id}) - SessionSupervisor.close_session(state.registry, state.module, session_id) - {:noreply, %{state | sessions: Map.delete(state.sessions, session_id)}} - else - {:noreply, state} - end - end - - def handle_info({:send_sampling_request, params, timeout}, state) do - request_id = ID.generate_request_id() - handle_sampling_request_send(request_id, params, timeout, state) - end - - def handle_info({:sampling_request_timeout, request_id}, state) do - handle_sampling_timeout(request_id, state) - end - - def handle_info({:send_roots_request, timeout}, state) do - request_id = ID.generate_request_id() - handle_roots_request_send(request_id, timeout, state) - end - - def handle_info({:roots_request_timeout, request_id}, state) do - handle_roots_timeout(request_id, state) - end - - def handle_info(event, %{module: module} = state) do - if Anubis.exported?(module, :handle_info, 2) do - case module.handle_info(event, state.frame) do - {:noreply, frame} -> {:noreply, %{state | frame: frame}} - {:noreply, frame, cont} -> {:noreply, %{state | frame: frame}, cont} - {:stop, reason, frame} -> {:stop, reason, %{state | frame: frame}} - end - else - {:noreply, state} - end - end - - @impl GenServer - def terminate(reason, %{module: module, server_info: server_info} = state) do - Logging.server_event("terminating", %{reason: reason, server_info: server_info}) - - Telemetry.execute( - Telemetry.event_server_terminate(), - %{system_time: System.system_time()}, - %{reason: reason, server_info: server_info} - ) - - if Anubis.exported?(module, :terminate, 2) do - module.terminate(reason, state.frame) - else - :ok - end - end - - @impl GenServer - def format_status(status) do - Map.new(status, fn - {:state, state} -> - {:state, format_state(state)} - - {:message, {:request, decoded, session_id, _ctx}} -> - {:message, {:request, decoded, session_id}} - - {:message, {:notification, decoded, session_id, _ctx}} -> - {:message, {:notification, decoded, session_id}} - - {:message, {:response, decoded, session_id, _ctx}} -> - {:message, {:response, decoded, session_id}} - - other -> - other - end) - end - - @non_printable_keys ~w(transport sessions expiry_timers server_requests)a - - defp format_state(state) do - pending_requests = format_pending_requests(state.server_requests) - sessions = format_sessions(state.sessions) - - state - |> Map.reject(fn {k, _} -> k in @non_printable_keys end) - |> Map.merge(%{ - transport: state.transport[:layer], - pending_requests: pending_requests, - active_sessions: sessions - }) - end - - defp format_pending_requests(requests) do - Enum.map(requests, fn {id, req} -> - %{id: id, method: req[:method], session_id: req[:session_id]} - end) - end - - defp format_sessions(sessions) do - sessions - |> Enum.map(fn {_id, {name, _}} -> Session.get(name) end) - |> Enum.reject(&is_nil/1) - end - - defguardp is_server_initialized(decoded, session) - when Message.is_initialize_lifecycle(decoded) or - Session.is_initialized(session) - - defp handle_single_request(decoded, session, state) do - cond do - Message.is_response(decoded) and server_request?(decoded["id"], state) -> - handle_server_request_response(decoded, state) - - Message.is_error(decoded) and server_request?(decoded["id"], state) -> - handle_server_request_error(decoded, state) - - Message.is_ping(decoded) -> - handle_server_ping(decoded, state) - - not is_server_initialized(decoded, session) -> - handle_server_not_initialized(state) - - Message.is_request(decoded) -> - handle_request(decoded, session, state) - - true -> - handle_invalid_request(state) - end - end - - defp handle_server_ping(%{"id" => request_id}, state) do - {:reply, {:ok, Message.build_response(%{}, request_id)}, state} - end - - defp handle_server_not_initialized(state) do - error = Error.protocol(:invalid_request, %{message: "Server not initialized"}) - - Logging.server_event( - "request_error", - %{error: error, reason: "not_initialized"}, - level: :warning - ) - - {:reply, {:ok, Error.build_json_rpc(error)}, state} - end - - defp handle_invalid_request(state) do - error = - Error.protocol(:invalid_request, %{ - message: "Expected request but got different message type" - }) - - {:reply, {:error, error}, state} - end - - # Request handling - - defp handle_request(%{"params" => params} = request, session, state) when Message.is_initialize(request) do - %{ - "clientInfo" => client_info, - "capabilities" => client_capabilities, - "protocolVersion" => requested_version - } = params - - {:ok, protocol_version, protocol_module} = - Anubis.Protocol.Registry.negotiate(requested_version, state.supported_versions) - - :ok = - Session.update_from_initialization( - session.name, - protocol_version, - client_info, - client_capabilities, - protocol_module: protocol_module - ) - - result = %{ - "protocolVersion" => protocol_version, - "serverInfo" => state.server_info, - "capabilities" => state.capabilities - } - - Logging.server_event("initializing", %{ - client_info: params["clientInfo"], - client_capabilities: params["capabilities"], - protocol_version: protocol_version - }) - - Telemetry.execute( - Telemetry.event_server_response(), - %{system_time: System.system_time()}, - %{method: "initialize", status: :success} - ) - - {:reply, {:ok, Message.build_response(result, request["id"])}, state} - end - - defp handle_request(%{"id" => request_id, "method" => "logging/setLevel"} = request, session, state) - when Server.is_supported_capability(state.capabilities, "logging") do - level = request["params"]["level"] - :ok = Session.set_log_level(session.name, level) - {:reply, {:ok, Message.build_response(%{}, request_id)}, state} - end - - defp handle_request(%{"id" => request_id, "method" => method} = request, session, state) do - Logging.server_event("handling_request", %{id: request_id, method: method}) - - :ok = Session.track_request(session.name, request_id, method) - - Telemetry.execute( - Telemetry.event_server_request(), - %{system_time: System.system_time()}, - %{id: request_id, method: method} - ) - - frame = - Frame.put_request(state.frame, %{ - id: request_id, - method: method, - params: request["params"] || %{} - }) - - server_request(request, %{state | frame: frame}) - end - - # Notification handling - - defp handle_notification(%{"method" => "notifications/initialized"}, session, %{module: module} = state) do - Logging.server_event("client_initialized", %{session_id: session.id}) - :ok = Session.mark_initialized(session.name) - - Logging.server_event("session_marked_initialized", %{ - session_id: session.id, - initialized: true - }) - - frame = %{state.frame | initialized: true} - - {:ok, frame} = - if Anubis.exported?(module, :init, 2), - do: module.init(session.client_info, frame), - else: {:ok, frame} - - {:noreply, %{state | frame: frame}} - end - - defp handle_notification(%{"method" => "notifications/cancelled"} = notification, session, state) do - params = notification["params"] || %{} - request_id = params["requestId"] - reason = Map.get(params, "reason", "cancelled") - - if Session.has_pending_request?(session.name, request_id) do - request_info = Session.complete_request(session.name, request_id) - - Logging.server_event("request_cancelled", %{ - session_id: session.id, - request_id: request_id, - reason: reason, - method: request_info[:method], - duration_ms: System.system_time(:millisecond) - request_info[:started_at] - }) - - Telemetry.execute( - Telemetry.event_server_notification(), - %{system_time: System.system_time()}, - %{method: "cancelled", session_id: session.id, request_id: request_id} - ) - - {:noreply, state} - else - Logging.server_event("cancellation_for_unknown_request", %{ - session_id: session.id, - request_id: request_id, - reason: reason - }) - - {:noreply, state} - end - end - - defp handle_notification(notification, _session, state) do - method = notification["method"] - - Logging.server_event("handling_notification", %{method: method}) - - Telemetry.execute( - Telemetry.event_server_notification(), - %{system_time: System.system_time()}, - %{method: method} - ) - - server_notification(notification, state) - end - - # Helper functions - - defp server_request(%{"id" => request_id, "method" => method} = request, %{module: module} = state) do - case module.handle_request(request, state.frame) do - {:reply, response, %Frame{} = frame} -> - Telemetry.execute( - Telemetry.event_server_response(), - %{system_time: System.system_time()}, - %{id: request_id, method: method, status: :success} - ) - - frame = Frame.clear_request(frame) - - {:reply, {:ok, Message.build_response(response, request_id)}, %{state | frame: frame}} - - {:noreply, %Frame{} = frame} -> - Telemetry.execute( - Telemetry.event_server_response(), - %{system_time: System.system_time()}, - %{id: request_id, method: method, status: :noreply} - ) - - frame = Frame.clear_request(frame) - {:reply, {:ok, nil}, %{state | frame: frame}} - - {:error, %Error{} = error, %Frame{} = frame} -> - Logging.server_event( - "request_error", - %{id: request_id, method: method, error: error}, - level: :warning - ) - - Telemetry.execute( - Telemetry.event_server_error(), - %{system_time: System.system_time()}, - %{id: request_id, method: method, error: error} - ) - - frame = Frame.clear_request(frame) - - {:reply, {:ok, Error.build_json_rpc(error, request_id)}, %{state | frame: frame}} - end - end - - defp server_notification(%{"method" => method} = notification, %{module: module} = state) do - case module.handle_notification(notification, state.frame) do - {:noreply, %Frame{} = frame} -> - {:noreply, %{state | frame: frame}} - - {:error, _error, %Frame{} = frame} -> - Logging.server_event( - "notification_handler_error", - %{method: method}, - level: :warning - ) - - {:noreply, %{state | frame: frame}} - end - end - - @spec maybe_attach_session(session_id :: String.t(), map, t) :: - {:ok, {session :: Session.t(), t}} - defp maybe_attach_session(session_id, context, %{sessions: sessions} = state) when is_map_key(sessions, session_id) do - {session_name, _ref} = sessions[session_id] - session = Session.get(session_name) - state = reset_session_expiry(session_id, state) - - {:ok, {session, %{state | frame: populate_frame(state.frame, session, context, state)}}} - end - - defp maybe_attach_session(session_id, context, %{sessions: sessions, registry: registry} = state) do - session_name = registry.server_session(state.module, session_id) - - case SessionSupervisor.create_session(registry, state.module, session_id) do - {:ok, pid} -> - ref = Process.monitor(pid) - - state = %{ - state - | sessions: Map.put(sessions, session_id, {session_name, ref}) - } - - state = reset_session_expiry(session_id, state) - - session = Session.get(session_name) - - {:ok, {session, %{state | frame: populate_frame(state.frame, session, context, state)}}} - - {:error, {:already_started, pid}} -> - ref = Process.monitor(pid) - - state = %{ - state - | sessions: Map.put(sessions, session_id, {session_name, ref}) - } - - state = reset_session_expiry(session_id, state) - - session = Session.get(session_name) - - {:ok, {session, %{state | frame: populate_frame(state.frame, session, context, state)}}} - - error -> - error - end - end - - defp populate_frame(frame, %Session{} = session, context, state) do - {assigns, context} = Map.pop(context, :assigns, %{}) - assigns = Map.merge(frame.assigns, assigns) - - frame - |> Frame.put_transport(context) - |> Frame.assign(assigns) - |> Frame.put_private(%{ - session_id: session.id, - client_info: session.client_info, - client_capabilities: session.client_capabilities, - protocol_version: session.protocol_version, - protocol_module: session.protocol_module, - server_registry: state.registry, - server_module: state.module - }) - end - - defp encode_notification(method, params) do - notification = Message.build_notification(method, params) - Logging.message("outgoing", "notification", nil, notification) - Message.encode_notification(notification) - end - - defp send_to_transport(nil, _data, _opts) do - {:error, Error.transport(:no_transport, %{message: "No transport configured"})} - end - - defp send_to_transport(%{layer: layer, name: name}, data, opts) do - with {:error, reason} <- layer.send_message(name, data, opts) do - {:error, Error.transport(:send_failure, %{original_reason: reason})} - end - end - - # Session expiry timer management - - defp schedule_session_expiry(session_id, timeout) do - Process.send_after(self(), {:session_expired, session_id}, timeout) - end - - defp reset_session_expiry(session_id, %{expiry_timers: timers, session_idle_timeout: timeout} = state) do - if timer = Map.get(timers, session_id), do: Process.cancel_timer(timer) - - timer = schedule_session_expiry(session_id, timeout) - %{state | expiry_timers: Map.put(timers, session_id, timer)} - end - - defp cancel_session_expiry(session_id, %{expiry_timers: timers} = state) do - if timer = Map.get(timers, session_id) do - Process.cancel_timer(timer) - %{state | expiry_timers: Map.delete(timers, session_id)} - else - state - end - end - - # Sampling request helpers - - defp handle_sampling_request_send(request_id, params, timeout, state) do - timer_ref = - Process.send_after(self(), {:sampling_request_timeout, request_id}, timeout) - - request_info = %{ - method: "sampling/createMessage", - session_id: state.frame.private.session_id, - timer_ref: timer_ref - } - - state = put_in(state.server_requests[request_id], request_info) - - with :ok <- validate_client_capability(state, "sampling"), - {:ok, request_data} <- - encode_request("sampling/createMessage", params, request_id), - :ok <- send_to_transport(state.transport, request_data, timeout: state.timeout) do - Logging.server_event("sent_sampling_request", %{request_id: request_id}) - {:noreply, state} - else - {:error, error} -> - Process.cancel_timer(timer_ref) - - state = %{ - state - | server_requests: Map.delete(state.server_requests, request_id) - } - - Logging.server_event( - "failed_send_sampling_request", - %{request_id: request_id, error: error}, - level: :error - ) - - {:noreply, state} - end - end - - defp validate_client_capability(%{frame: frame} = state, capability) do - current_session = Frame.get_mcp_session_id(frame) - - session_name = - Enum.find_value(state.sessions, fn {id, {name, _ref}} -> - if id == current_session, do: name - end) - - session = Session.get(session_name) - - if Map.has_key?(session.client_capabilities || %{}, capability) do - :ok - else - {:error, "No session initialzied for sending sampling request"} - end - end - - defp handle_sampling_timeout(request_id, state) do - case Map.pop(state.server_requests, request_id) do - {nil, _} -> - {:noreply, state} - - {_request_info, updated_requests} -> - Logging.server_event("sampling_request_timeout", %{request_id: request_id}, level: :warning) - - {:noreply, %{state | server_requests: updated_requests}} - end - end - - defp encode_request(method, params, request_id) do - request = %{ - "method" => method, - "params" => params - } - - Logging.message("outgoing", "request", request_id, request) - Message.encode_request(request, request_id) - end - - defp server_request?(request_id, %{server_requests: requests}) when is_binary(request_id) do - Map.has_key?(requests, request_id) - end - - defp server_request?(_, _), do: false - - defp handle_server_request_response(%{"id" => request_id, "result" => result}, state) do - {request_info, updated_requests} = Map.pop(state.server_requests, request_id) - Process.cancel_timer(request_info.timer_ref) - - state = %{state | server_requests: updated_requests} - - case request_info.method do - "sampling/createMessage" -> - handle_sampling(result, request_id, state) - - "roots/list" -> - handle_roots(result["roots"] || [], request_id, state) - - _ -> - {:noreply, state} - end - end - - defp handle_server_request_error(%{"id" => request_id, "error" => error}, state) do - {request_info, updated_requests} = Map.pop(state.server_requests, request_id) - Process.cancel_timer(request_info.timer_ref) - - state = %{state | server_requests: updated_requests} - - Logging.server_event( - "server_request_error", - %{ - request_id: request_id, - method: request_info.method, - error: error - }, - level: :error - ) - - {:noreply, state} - end - - defp handle_sampling(result, request_id, %{module: module, frame: frame} = state) do - case module.handle_sampling(result, request_id, frame) do - {:noreply, new_frame} -> - {:noreply, %{state | frame: new_frame}} - - {:stop, reason, new_frame} -> - {:stop, reason, %{state | frame: new_frame}} - end - end - - # Roots request helpers - - defp handle_roots_request_send(request_id, timeout, state) do - timer_ref = - Process.send_after(self(), {:roots_request_timeout, request_id}, timeout) - - request_info = %{ - id: request_id, - method: "roots/list", - session_id: state.frame.private.session_id, - timer_ref: timer_ref - } - - state = put_in(state.server_requests[request_id], request_info) - - with :ok <- validate_client_capability(state, "roots"), - {:ok, request_data} <- encode_request("roots/list", %{}, request_id), - :ok <- send_to_transport(state.transport, request_data, timeout: state.timeout) do - Logging.server_event("sent_roots_request", %{request_id: request_id}) - {:noreply, state} - else - {:error, error} -> - Process.cancel_timer(timer_ref) - - state = %{ - state - | server_requests: Map.delete(state.server_requests, request_id) - } - - Logging.server_event( - "failed_send_roots_request", - %{request_id: request_id, error: error}, - level: :error - ) - - {:noreply, state} - end - end - - defp handle_roots_timeout(request_id, state) when is_binary(request_id) do - state.server_requests - |> Map.pop(request_id) - |> handle_roots_timeout(state) - end - - defp handle_roots_timeout({nil, _}, state), do: {:noreply, state} - - defp handle_roots_timeout({%{id: request_id}, requests}, state) do - with {:ok, notification} <- - encode_notification("notifications/cancelled", %{ - "requestId" => request_id, - "reason" => "timeout" - }), - :ok <- send_to_transport(state.transport, notification, timeout: state.timeout) do - Logging.server_event( - "roots_request_timeout_cancelled", - %{request_id: request_id} - ) - end - - Logging.server_event("roots_request_timeout", %{request_id: request_id}, level: :warning) - - {:noreply, %{state | server_requests: requests}} - end - - defp handle_roots(roots, request_id, %{module: module} = state) do - case module.handle_roots(roots, request_id, state.frame) do - {:noreply, new_frame} -> - {:noreply, %{state | frame: new_frame}} - - {:stop, reason, new_frame} -> - {:stop, reason, %{state | frame: new_frame}} - end - end -end diff --git a/lib/anubis/server/context.ex b/lib/anubis/server/context.ex new file mode 100644 index 00000000..30fafad2 --- /dev/null +++ b/lib/anubis/server/context.ex @@ -0,0 +1,23 @@ +defmodule Anubis.Server.Context do + @moduledoc """ + Read-only session and request context, set by the SDK before each callback. + + The Session process builds a fresh Context before every user callback invocation. + Mutations have no lasting effect — the Session always overwrites it. + + For STDIO transport, `headers` is empty and `remote_ip` is nil. + For HTTP transport, headers are normalized to lowercase string keys. + """ + + @type t :: %__MODULE__{ + session_id: String.t() | nil, + client_info: map() | nil, + headers: %{String.t() => String.t()}, + remote_ip: :inet.ip_address() | nil + } + + defstruct session_id: nil, + client_info: nil, + headers: %{}, + remote_ip: nil +end diff --git a/lib/anubis/server/frame.ex b/lib/anubis/server/frame.ex index 2f1e6bd0..284d2dc9 100644 --- a/lib/anubis/server/frame.ex +++ b/lib/anubis/server/frame.ex @@ -1,59 +1,28 @@ defmodule Anubis.Server.Frame do @moduledoc """ - The Anubis Frame. - - This module defines a struct and functions for working with - MCP server state throughout the request/response lifecycle. + The Anubis Frame — pure user state + read-only context. ## User fields - These fields contain user-controlled data: - * `assigns` - shared user data as a map. For HTTP transports, this inherits - from `Plug.Conn.assigns`. Users are responsible for populating authentication - data through their Plug pipeline before it reaches the MCP server. - - ## Transport fields - - These fields contain transport-specific context. The structure varies by transport type: - - ### HTTP transport (when `transport.type == :http`) + from `Plug.Conn.assigns`. - * `req_headers` - the request headers as a list, example: `[{"content-type", "application/json"}]`. - All header names are downcased. - * `query_params` - the request query params as a map, example: `%{"session" => "abc123"}`. - Returns `nil` if query params were not fetched by the Plug pipeline. - * `remote_ip` - the IP of the client, example: `{151, 236, 219, 228}`. - This field is set by the transport layer. - * `scheme` - the request scheme as an atom, example: `:https` - * `host` - the requested host as a binary, example: `"api.example.com"` - * `port` - the requested port as an integer, example: `443` - * `request_path` - the requested path, example: `"/mcp"` + ## Component maps - ### STDIO transport (when `transport.type == :stdio`) + Runtime-registered components are stored in typed maps keyed by name/URI: - * `env` - environment variables as a map, example: `%{"USER" => "alice", "HOME" => "/home/alice"}` - * `pid` - the OS process ID as a string, example: `"12345"` + * `tools` - `%{name => %Tool{}}` + * `resources` - `%{uri => %Resource{}}` + * `prompts` - `%{name => %Prompt{}}` + * `resource_templates` - `%{name => %Resource{uri_template: ...}}` - ## MCP protocol fields + ## Pagination - These fields contain MCP-specific data: + * `pagination_limit` - optional limit for listing operations - * `request` - the current MCP request being processed, with fields: - * `id` - the request ID for correlation - * `method` - the MCP method being called, example: `"tools/call"` - * `params` - the raw request parameters (before validation) - * `initialized` - boolean indicating if the MCP session has been initialized + ## Context - ## Private fields - - These fields are reserved for framework usage: - - * `private` - shared framework data as a map. Contains MCP session context: - * `session_id` - unique identifier for the current client session being handled - * `client_info` - client information from initialization, example: `%{"name" => "my-client", "version" => "1.0.0"}` - * `client_capabilities` - negotiated client capabilities - * `protocol_version` - active MCP protocol version, example: `"2025-03-26"` + * `context` - read-only `%Context{}`, refreshed by Session before each callback """ alias Anubis.Server.Component @@ -61,58 +30,27 @@ defmodule Anubis.Server.Frame do alias Anubis.Server.Component.Resource alias Anubis.Server.Component.Schema alias Anubis.Server.Component.Tool + alias Anubis.Server.Context @type server_component_t :: Tool.t() | Resource.t() | Prompt.t() - @type private_t :: %{ - optional(:session_id) => String.t(), - optional(:client_info) => map(), - optional(:client_capabilities) => map(), - optional(:protocol_version) => String.t(), - optional(:server_module) => module(), - optional(:server_registry) => module(), - optional(:pagination_limit) => non_neg_integer(), - optional(:__mcp_components__) => list(server_component_t) - } - - @type request_t :: %{ - id: String.t(), - method: String.t(), - params: map() - } - - @type http_t :: %{ - type: :http, - req_headers: [{String.t(), String.t()}], - query_params: %{optional(String.t()) => String.t()} | nil, - remote_ip: term, - scheme: :http | :https, - host: String.t(), - port: non_neg_integer, - request_path: String.t() - } - - @type stdio_t :: %{ - type: :stdio, - os_pid: non_neg_integer, - env: map - } - - @type transport_t :: http_t | stdio_t - @type t :: %__MODULE__{ - assigns: Enumerable.t(), - initialized: boolean, - private: private_t, - request: request_t | nil, - transport: transport_t + assigns: map(), + tools: %{optional(String.t()) => Tool.t()}, + resources: %{optional(String.t()) => Resource.t()}, + prompts: %{optional(String.t()) => Prompt.t()}, + resource_templates: %{optional(String.t()) => Resource.t()}, + pagination_limit: non_neg_integer() | nil, + context: Context.t() } defstruct assigns: %{}, - initialized: false, - private: %{}, - request: nil, - transport: %{} + tools: %{}, + resources: %{}, + prompts: %{}, + resource_templates: %{}, + pagination_limit: nil, + context: %Context{} @doc """ Creates a new frame with optional initial assigns. @@ -120,13 +58,13 @@ defmodule Anubis.Server.Frame do ## Examples iex> Frame.new() - %Frame{assigns: %{}, initialized: false} + %Frame{assigns: %{}} iex> Frame.new(%{user: "alice"}) - %Frame{assigns: %{user: "alice"}, initialized: false} + %Frame{assigns: %{user: "alice"}} """ - @spec new :: t - @spec new(assigns :: Enumerable.t()) :: t + @spec new :: t() + @spec new(assigns :: map()) :: t() def new(assigns \\ %{}), do: struct(__MODULE__, assigns: assigns) @doc """ @@ -134,17 +72,12 @@ defmodule Anubis.Server.Frame do ## Examples - # Single assignment frame = Frame.assign(frame, :status, :active) - - # Multiple assignments via map frame = Frame.assign(frame, %{status: :active, count: 5}) - - # Multiple assignments via keyword list frame = Frame.assign(frame, status: :active, count: 5) """ - @spec assign(t, Enumerable.t()) :: t - @spec assign(t, key :: atom, value :: any) :: t + @spec assign(t(), Enumerable.t()) :: t() + @spec assign(t(), key :: atom(), value :: any()) :: t() def assign(%__MODULE__{} = frame, assigns) when is_map(assigns) or is_list(assigns) do Enum.reduce(assigns, frame, fn {key, value}, frame -> assign(frame, key, value) @@ -158,20 +91,13 @@ defmodule Anubis.Server.Frame do @doc """ Assigns a value to the frame only if the key doesn't already exist. - The value is computed lazily using the provided function, which is only - called if the key is not present in assigns. + The value is computed lazily using the provided function. ## Examples - # Only assigns if :timestamp doesn't exist frame = Frame.assign_new(frame, :timestamp, fn -> DateTime.utc_now() end) - - # Function is not called if key exists - frame = frame |> Frame.assign(:count, 5) - |> Frame.assign_new(:count, fn -> expensive_computation() end) - # count remains 5 """ - @spec assign_new(t, key :: atom, value_fun :: (-> term)) :: t + @spec assign_new(t(), key :: atom(), value_fun :: (-> term())) :: t() def assign_new(%__MODULE__{} = frame, key, fun) when is_atom(key) and is_function(fun, 0) do case frame.assigns do %{^key => _} -> frame @@ -179,261 +105,30 @@ defmodule Anubis.Server.Frame do end end - @doc """ - Sets or updates private session data in the frame. - - Private data is used for framework-internal session context that persists - across requests, similar to Plug.Conn.private. - - ## Examples - - # Set single private value - frame = Frame.put_private(frame, :session_id, "abc123") - - # Set multiple private values - frame = Frame.put_private(frame, %{ - session_id: "abc123", - client_info: %{name: "my-client", version: "1.0.0"} - }) - """ - @spec put_private(t, atom, any) :: t - @spec put_private(t, Enumerable.t()) :: t - def put_private(%__MODULE__{} = frame, key, value) when is_atom(key) do - %{frame | private: Map.put(frame.private, key, value)} - end - - def put_private(%__MODULE__{} = frame, private) when is_map(private) or is_list(private) do - Enum.reduce(private, frame, fn {key, value}, frame -> - put_private(frame, key, value) - end) - end - - @doc """ - Sets or updates transport data in the frame. - - Check `transport_t()` for reference. - - ## Examples - - # Set single transport value - frame = Frame.put_transport(frame, :session_id, "abc123") - - # Set multiple transport values - frame = Frame.put_transport(frame, %{ - session_id: "abc123", - client_info: %{name: "my-client", version: "1.0.0"} - }) - """ - @spec put_transport(t, atom, any) :: t - @spec put_transport(t, Enumerable.t()) :: t - def put_transport(%__MODULE__{} = frame, key, value) when is_atom(key) do - %{frame | transport: Map.put(frame.transport, key, value)} - end - - def put_transport(%__MODULE__{} = frame, transport) when is_map(transport) or is_list(transport) do - Enum.reduce(transport, frame, fn {key, value}, frame -> - put_transport(frame, key, value) - end) - end - - @doc """ - Sets the current request being processed. - - The request includes the request ID, method, and raw parameters before validation. - - ## Examples - - frame = Frame.put_request(frame, %{ - id: "req_123", - method: "tools/call", - params: %{"name" => "calculator", "arguments" => %{}} - }) - """ - @spec put_request(t, map) :: t - def put_request(%__MODULE__{} = frame, request) when is_map(request) do - %{frame | request: request} - end - @doc """ Sets the pagination limit for listing operations. - This limit is used by handlers when returning lists of tools, prompts, or resources - to control the maximum number of items returned in a single response. When the limit - is set and the total number of items exceeds it, the response will include a - `nextCursor` field for pagination. - ## Examples - # Set pagination limit to 10 items per page frame = Frame.put_pagination_limit(frame, 10) - - # The limit is stored in private data - frame.private.pagination_limit + frame.pagination_limit # => 10 """ - @spec put_pagination_limit(t, non_neg_integer) :: t + @spec put_pagination_limit(t(), non_neg_integer()) :: t() def put_pagination_limit(%__MODULE__{} = frame, limit) when limit > 0 do - put_private(frame, %{pagination_limit: limit}) - end - - @doc """ - Clears the current request from the frame. - - This should be called after processing a request to ensure the frame doesn't - retain stale request data. - - ## Examples - - frame = Frame.clear_request(frame) - """ - @spec clear_request(t) :: t - def clear_request(%__MODULE__{} = frame) do - %{frame | request: nil} - end - - @doc """ - Clears all session-specific private data from the frame. - - This should be called when a session ends to ensure the frame doesn't - retain stale session data. - - ## Examples - - frame = Frame.clear_session(frame) - """ - @spec clear_session(t) :: t - def clear_session(%__MODULE__{} = frame) do - %{frame | private: %{}} + %{frame | pagination_limit: limit} end @doc """ - Gets the MCP session ID from the frame's private data. - - ## Examples - - session_id = Frame.get_mcp_session_id(frame) - # => "session_abc123" + Registers a tool definition at runtime. """ - @spec get_mcp_session_id(t) :: String.t() | nil - def get_mcp_session_id(%__MODULE__{} = frame) do - Map.get(frame.private, :session_id) - end - - @doc """ - Gets the client info from the frame's private data. - - ## Examples - - client_info = Frame.get_client_info(frame) - # => %{"name" => "my-client", "version" => "1.0.0"} - """ - @spec get_client_info(t) :: map() | nil - def get_client_info(%__MODULE__{} = frame) do - Map.get(frame.private, :client_info) - end - - @doc """ - Gets the client capabilities from the frame's private data. - - ## Examples - - capabilities = Frame.get_client_capabilities(frame) - # => %{"tools" => %{}, "resources" => %{}} - """ - @spec get_client_capabilities(t) :: map() | nil - def get_client_capabilities(%__MODULE__{} = frame) do - Map.get(frame.private, :client_capabilities) - end - - @doc """ - Gets the protocol version from the frame's private data. - - ## Examples - - version = Frame.get_protocol_version(frame) - # => "2025-03-26" - """ - @spec get_protocol_version(t) :: String.t() | nil - def get_protocol_version(%__MODULE__{} = frame) do - Map.get(frame.private, :protocol_version) - end - - @doc """ - Gets the negotiated protocol module from the frame's private data. - - Returns the module implementing `Anubis.Protocol.Behaviour` for the - negotiated protocol version, or nil if not yet negotiated. - - ## Examples - - mod = Frame.get_protocol_module(frame) - # => Anubis.Protocol.V2025_03_26 - """ - @spec get_protocol_module(t) :: module() | nil - def get_protocol_module(%__MODULE__{} = frame) do - Map.get(frame.private, :protocol_module) - end - - @doc """ - Gets a request header value from HTTP transport. - - Returns the first value for the header, or nil if the transport - is not HTTP or the header is not present. - - ## Examples - - # HTTP transport - auth_header = Frame.get_req_header(frame, "authorization") - # => "Bearer token123" - - # Non-HTTP transport or missing header - auth_header = Frame.get_req_header(frame, "authorization") - # => nil - """ - @spec get_req_header(t, String.t()) :: String.t() | nil - def get_req_header(%__MODULE__{transport: %{type: :http, req_headers: headers}}, name) when is_binary(name) do - case List.keyfind(headers, String.downcase(name), 0) do - {_, value} -> value - nil -> nil - end - end - - def get_req_header(%__MODULE__{}, _name), do: nil - - @doc """ - Gets a query parameter value from HTTP transport. - - Returns the parameter value, or nil if the transport is not HTTP, - query params weren't fetched, or the parameter doesn't exist. - - ## Examples - - # HTTP transport with query params - session = Frame.get_query_param(frame, "session") - # => "abc123" - - # Missing parameter or non-HTTP transport - missing = Frame.get_query_param(frame, "nonexistent") - # => nil - """ - @spec get_query_param(t, String.t()) :: String.t() | nil - def get_query_param(%__MODULE__{transport: %{type: :http, query_params: params}}, key) - when is_map(params) and is_binary(key) do - Map.get(params, key) - end - - def get_query_param(%__MODULE__{}, _key), do: nil - - @doc """ - Registers a tool definition. - """ - @spec register_tool(t, String.t(), list(tool_opt)) :: t + @spec register_tool(t(), String.t(), list(tool_opt)) :: t() when tool_opt: {:description, String.t() | nil} - | {:input_schema, map | nil} - | {:output_schema, map | nil} + | {:input_schema, map() | nil} + | {:output_schema, map() | nil} | {:title, String.t() | nil} - | {:annotations, map | nil} + | {:annotations, map() | nil} def register_tool(%__MODULE__{} = frame, name, opts) when is_binary(name) do input_schema = Schema.normalize(opts[:input_schema] || %{}) raw_schema = Component.__clean_schema_for_peri__(input_schema) @@ -450,7 +145,7 @@ defmodule Anubis.Server.Frame do annotations = opts[:annotations] title = annotations[:title] || annotations["title"] || opts[:title] || name - update_components(frame, %Tool{ + tool = %Tool{ name: name, description: opts[:description], input_schema: Schema.to_json_schema(input_schema), @@ -459,27 +154,31 @@ defmodule Anubis.Server.Frame do title: title, validate_input: validate_input, validate_output: validate_output - }) + } + + %{frame | tools: Map.put(frame.tools, name, tool)} end @doc """ - Registers a prompt definition. + Registers a prompt definition at runtime. """ - @spec register_prompt(t, String.t(), list(prompt_opt)) :: t - when prompt_opt: {:description, String.t() | nil} | {:arguments, map | nil} | {:title, String.t() | nil} + @spec register_prompt(t(), String.t(), list(prompt_opt)) :: t() + when prompt_opt: {:description, String.t() | nil} | {:arguments, map() | nil} | {:title, String.t() | nil} def register_prompt(%__MODULE__{} = frame, name, opts) when is_binary(name) do arguments = Schema.normalize(opts[:arguments] || %{}) raw_schema = Component.__clean_schema_for_peri__(arguments) validate_input = fn params -> Peri.validate(raw_schema, params) end title = opts[:title] || name - update_components(frame, %Prompt{ + prompt = %Prompt{ name: name, title: title, description: opts[:description], arguments: Schema.to_prompt_arguments(arguments), validate_input: validate_input - }) + } + + %{frame | prompts: Map.put(frame.prompts, name, prompt)} end @doc """ @@ -487,7 +186,7 @@ defmodule Anubis.Server.Frame do For parameterized resources, use `register_resource_template/3` instead. """ - @spec register_resource(t, String.t(), list(resource_opt)) :: t + @spec register_resource(t(), String.t(), list(resource_opt)) :: t() when resource_opt: {:title, String.t() | nil} | {:name, String.t() | nil} @@ -496,20 +195,20 @@ defmodule Anubis.Server.Frame do def register_resource(%__MODULE__{} = frame, uri, opts) when is_binary(uri) do name = opts[:name] || Path.basename(uri) - update_components(frame, %Resource{ + resource = %Resource{ uri: uri, title: opts[:title] || name, name: name, description: opts[:description], mime_type: opts[:mime_type] || "text/plain" - }) + } + + %{frame | resources: Map.put(frame.resources, uri, resource)} end @doc """ Registers a resource template definition using a URI template (RFC 6570). - URI templates allow parameterized resources like `file:///{path}` or `db:///{table}/{id}`. - ## Examples frame = Frame.register_resource_template(frame, "file:///{path}", @@ -518,106 +217,120 @@ defmodule Anubis.Server.Frame do description: "Access files in the project directory" ) """ - @spec register_resource_template(t, String.t(), list(resource_template_opt)) :: t + @spec register_resource_template(t(), String.t(), list(resource_template_opt)) :: t() when resource_template_opt: {:title, String.t() | nil} | {:name, String.t()} | {:description, String.t() | nil} | {:mime_type, String.t() | nil} def register_resource_template(%__MODULE__{} = frame, uri_template, opts) when is_binary(uri_template) do - # name is required as it serves as a semantic identifier for the template. - # Unlike static resources, templates like "file:///{path}" cannot derive meaningful names. name = Keyword.fetch!(opts, :name) - update_components(frame, %Resource{ + resource = %Resource{ uri_template: uri_template, title: opts[:title] || name, name: name, description: opts[:description], mime_type: opts[:mime_type] || "text/plain" - }) + } + + %{frame | resource_templates: Map.put(frame.resource_templates, name, resource)} end - @doc "Clears all current registered components (tools, resources, prompts)" - @spec clear_components(t) :: t + @doc "Clears all runtime-registered components" + @spec clear_components(t()) :: t() def clear_components(%__MODULE__{} = frame) do - put_in(frame, [Access.key!(:private), :__mcp_components__], []) + %{frame | tools: %{}, resources: %{}, prompts: %{}, resource_templates: %{}} end - @doc "Retrieves all current registered components (tools, resources, prompts)" - @spec get_components(t) :: list(server_component_t) + @doc "Retrieves all runtime-registered components as a flat list" + @spec get_components(t()) :: list(server_component_t()) def get_components(%__MODULE__{} = frame) do - Map.get(frame.private, :__mcp_components__, []) + Map.values(frame.tools) ++ + Map.values(frame.resources) ++ + Map.values(frame.prompts) ++ + Map.values(frame.resource_templates) end @doc false - @spec get_tools(t) :: list(Tool.t()) - def get_tools(%__MODULE__{} = frame) do - frame - |> get_components() - |> Enum.filter(&match?(%Tool{}, &1)) - end + @spec get_tools(t()) :: list(Tool.t()) + def get_tools(%__MODULE__{} = frame), do: Map.values(frame.tools) @doc false - @spec get_prompts(t) :: list(Prompt.t()) - def get_prompts(%__MODULE__{} = frame) do - frame - |> get_components() - |> Enum.filter(&match?(%Prompt{}, &1)) - end + @spec get_prompts(t()) :: list(Prompt.t()) + def get_prompts(%__MODULE__{} = frame), do: Map.values(frame.prompts) @doc false - @spec get_resources(t) :: list(Resource.t()) + @spec get_resources(t()) :: list(Resource.t()) def get_resources(%__MODULE__{} = frame) do - frame - |> get_components() - |> Enum.filter(&match?(%Resource{}, &1)) + Map.values(frame.resources) ++ Map.values(frame.resource_templates) end @doc false - @spec get_component(t, name :: String.t()) :: server_component_t | nil + @spec get_component(t(), name :: String.t()) :: server_component_t() | nil def get_component(%__MODULE__{} = frame, name) do - frame - |> get_components() - |> Enum.find(&(&1.name == name)) + frame.tools[name] || + frame.prompts[name] || + frame.resource_templates[name] || + Enum.find(Map.values(frame.resources), &(&1.name == name)) end - # Private helpers + @doc """ + Serializes Frame for persistent storage. + + Only `assigns` and `pagination_limit` are persisted. The following fields are + **runtime-only** and excluded from serialization: + + * `tools` — runtime-registered tool definitions (includes validator functions) + * `resources` — runtime-registered resource definitions + * `prompts` — runtime-registered prompt definitions + * `resource_templates` — runtime-registered resource template definitions + * `context` — rebuilt by Session before each callback invocation - defp update_components(frame, component) do - components = [component | get_components(frame)] - put_private(frame, :__mcp_components__, Enum.uniq_by(components, &unique_component/1)) + Compile-time components (registered via the `component` macro) are always + available from the server module and do not need persistence. + """ + @spec to_saved(t()) :: map() + def to_saved(%__MODULE__{} = frame) do + %{ + "assigns" => frame.assigns, + "pagination_limit" => frame.pagination_limit + } end - defp unique_component(%struct{name: name}) do - {struct, name} + @doc """ + Reconstructs Frame from a previously saved map. + + Only `assigns` and `pagination_limit` are restored. Runtime-only fields (`tools`, + `resources`, `prompts`, `resource_templates`) are initialized empty — their validator + functions are not serializable. `context` is left as the default struct and will be + set by Session before each callback invocation. + """ + @spec from_saved(map()) :: t() + def from_saved(map) when is_map(map) do + %__MODULE__{ + assigns: Map.get(map, "assigns", %{}), + pagination_limit: Map.get(map, "pagination_limit") + } end + + def from_saved(_), do: %__MODULE__{} end defimpl Inspect, for: Anubis.Server.Frame do import Inspect.Algebra def inspect(frame, opts) do - components = frame.private[:__mcp_components__] || [] - tools_count = Enum.count(components, &match?(%Anubis.Server.Component.Tool{}, &1)) - - resources_count = - Enum.count(components, &match?(%Anubis.Server.Component.Resource{}, &1)) - - prompts_count = Enum.count(components, &match?(%Anubis.Server.Component.Prompt{}, &1)) - info = [ assigns: frame.assigns, - initialized: frame.initialized, - tools: tools_count, - resources: resources_count, - prompts: prompts_count + tools: map_size(frame.tools), + resources: map_size(frame.resources), + prompts: map_size(frame.prompts), + resource_templates: map_size(frame.resource_templates) ] - info = if frame.request, do: [{:request, frame.request.method} | info], else: info - info = - if session_id = frame.private[:session_id], + if session_id = frame.context.session_id, do: [{:session_id, session_id} | info], else: info diff --git a/lib/anubis/server/handlers.ex b/lib/anubis/server/handlers.ex index c3eb0de1..f820116e 100644 --- a/lib/anubis/server/handlers.ex +++ b/lib/anubis/server/handlers.ex @@ -52,15 +52,27 @@ defmodule Anubis.Server.Handlers do end def get_server_resources(module, frame) do - (module.__components__(:resource) ++ Frame.get_resources(frame)) - |> Enum.reject(& &1.uri_template) - |> Enum.sort_by(& &1.name) + compile_time = + :resource + |> module.__components__() + |> Enum.reject(& &1.uri_template) + + runtime = + Map.values(frame.resources) + + Enum.sort_by(compile_time ++ runtime, & &1.name) end def get_server_resource_templates(module, frame) do - (module.__components__(:resource) ++ Frame.get_resources(frame)) - |> Enum.filter(& &1.uri_template) - |> Enum.sort_by(& &1.name) + compile_time = + :resource + |> module.__components__() + |> Enum.filter(& &1.uri_template) + + runtime = + Map.values(frame.resource_templates) + + Enum.sort_by(compile_time ++ runtime, & &1.name) end @spec maybe_paginate(map, list(struct), non_neg_integer | nil) :: diff --git a/lib/anubis/server/handlers/prompts.ex b/lib/anubis/server/handlers/prompts.ex index fc5f974e..54f58b5d 100644 --- a/lib/anubis/server/handlers/prompts.ex +++ b/lib/anubis/server/handlers/prompts.ex @@ -12,7 +12,7 @@ defmodule Anubis.Server.Handlers.Prompts do {:reply, map(), Frame.t()} | {:error, Error.t(), Frame.t()} def handle_list(request, frame, server_module) do prompts = Handlers.get_server_prompts(server_module, frame) - limit = frame.private[:pagination_limit] + limit = frame.pagination_limit {prompts, cursor} = Handlers.maybe_paginate(request, prompts, limit) {:reply, diff --git a/lib/anubis/server/handlers/resources.ex b/lib/anubis/server/handlers/resources.ex index 37c7bb2b..4f5f17e7 100644 --- a/lib/anubis/server/handlers/resources.ex +++ b/lib/anubis/server/handlers/resources.ex @@ -11,7 +11,7 @@ defmodule Anubis.Server.Handlers.Resources do {:reply, map(), Frame.t()} | {:error, Error.t(), Frame.t()} def handle_list(request, frame, server_module) do resources = Handlers.get_server_resources(server_module, frame) - limit = frame.private[:pagination_limit] + limit = frame.pagination_limit {resources, cursor} = Handlers.maybe_paginate(request, resources, limit) {:reply, @@ -25,7 +25,7 @@ defmodule Anubis.Server.Handlers.Resources do {:reply, map(), Frame.t()} | {:error, Error.t(), Frame.t()} def handle_templates_list(request, frame, server_module) do templates = Handlers.get_server_resource_templates(server_module, frame) - limit = frame.private[:pagination_limit] + limit = frame.pagination_limit {templates, cursor} = Handlers.maybe_paginate(request, templates, limit) {:reply, diff --git a/lib/anubis/server/handlers/tools.ex b/lib/anubis/server/handlers/tools.ex index 3232b064..3cf9b382 100644 --- a/lib/anubis/server/handlers/tools.ex +++ b/lib/anubis/server/handlers/tools.ex @@ -12,7 +12,7 @@ defmodule Anubis.Server.Handlers.Tools do {:reply, map(), Frame.t()} | {:error, Error.t(), Frame.t()} def handle_list(request, frame, server_module) do tools = Handlers.get_server_tools(server_module, frame) - limit = frame.private[:pagination_limit] + limit = frame.pagination_limit {tools, cursor} = Handlers.maybe_paginate(request, tools, limit) {:reply, diff --git a/lib/anubis/server/registry.ex b/lib/anubis/server/registry.ex index 99f60600..2dd5111e 100644 --- a/lib/anubis/server/registry.ex +++ b/lib/anubis/server/registry.ex @@ -1,91 +1,46 @@ defmodule Anubis.Server.Registry do - @moduledoc false + @moduledoc """ + Behaviour for pluggable session registries and deterministic naming utilities. - def child_spec(_) do - Registry.child_spec(keys: :unique, name: __MODULE__) - end + The registry is responsible for mapping session IDs to PIDs. Different transports + have different needs: - @doc """ - Returns a via tuple for naming a server process. - """ - @spec server(server_module :: module()) :: GenServer.name() - def server(module) do - {:via, Registry, {__MODULE__, {:server, module}}} - end + - STDIO: single session, no registry needed (`Registry.None`) + - HTTP: multiple sessions, need lookup by session ID (`Registry.Local`) - @spec task_supervisor(server_module :: module()) :: GenServer.name() - def task_supervisor(module) when is_atom(module) do - {:via, Registry, {__MODULE__, {:task_supervisor, module}}} - end + ## Naming Utilities - @doc """ - Returns a via tuple for naming a server session process. + The module also provides deterministic atom naming for internal processes. + These are safe because server modules are compile-time bounded. """ - @spec server_session(server_module :: module(), session_id :: String.t()) :: - GenServer.name() - def server_session(server, session_id) do - {:via, Registry, {__MODULE__, {:session, server, session_id}}} - end - @doc """ - Returns a via tuple for naming a transport process. - """ - @spec transport(server_module :: module(), transport_type :: atom()) :: - GenServer.name() - def transport(module, type) when is_atom(module) do - {:via, Registry, {__MODULE__, {:transport, module, type}}} - end + @type session_id :: String.t() - @doc """ - Returns a via tuple for naming a supervisor process. - """ - @spec supervisor(kind :: atom(), server_module :: module()) :: GenServer.name() - def supervisor(kind \\ :supervisor, module) do - {:via, Registry, {__MODULE__, {kind, module}}} - end + @callback child_spec(keyword()) :: Supervisor.child_spec() | :ignore + @callback register_session(name :: term(), session_id(), pid()) :: :ok | {:error, term()} + @callback lookup_session(name :: term(), session_id()) :: {:ok, pid()} | {:error, :not_found} + @callback unregister_session(name :: term(), session_id()) :: :ok - @doc """ - Gets the PID of a session-specific server. - """ - @spec whereis_server_session(server_module :: module(), session_id :: String.t()) :: - pid | nil - def whereis_server_session(module, session_id) do - case Registry.lookup(__MODULE__, {:session, module, session_id}) do - [{pid, _}] -> pid - [] -> nil - end - end + # Deterministic atom naming for internal processes - @doc """ - Gets the PID of a supervisor process. - """ - @spec whereis_supervisor(atom(), module()) :: pid() | nil - def whereis_supervisor(server, kind \\ :supervisor) when is_atom(server) do - case Registry.lookup(__MODULE__, {kind, server}) do - [{pid, _}] -> pid - [] -> nil - end - end + @spec transport_name(module(), atom()) :: atom() + def transport_name(server, type), do: :"Anubis.#{server}.transport.#{type}" - @doc """ - Gets the PID of a registered server. - """ - @spec whereis_server(module()) :: pid | nil - def whereis_server(module) when is_atom(module) do - case Registry.lookup(__MODULE__, {:server, module}) do - [{pid, _}] -> pid - [] -> nil - end - end + @spec task_supervisor_name(module()) :: atom() + def task_supervisor_name(server), do: :"Anubis.#{server}.task_supervisor" - @doc """ - Gets the PID of a registered transport. - """ - @spec whereis_transport(module(), atom()) :: pid | nil - def whereis_transport(module, type) when is_atom(module) and is_atom(type) do - case Registry.lookup(__MODULE__, {:transport, module, type}) do - [{pid, _}] -> pid - [] -> nil - end - end + @spec session_supervisor_name(module()) :: atom() + def session_supervisor_name(server), do: :"Anubis.#{server}.session_supervisor" + + @spec supervisor_name(module()) :: atom() + def supervisor_name(server), do: :"Anubis.#{server}.supervisor" + + @spec session_name(module(), String.t()) :: atom() + def session_name(server, session_id), do: :"Anubis.#{server}.session.#{session_id}" + + @spec stdio_session_name(module()) :: atom() + def stdio_session_name(server), do: :"Anubis.#{server}.session.stdio" + + @spec registry_name(module()) :: atom() + def registry_name(server), do: :"Anubis.#{server}.registry" end diff --git a/lib/anubis/server/registry/adapter.ex b/lib/anubis/server/registry/adapter.ex deleted file mode 100644 index 89168891..00000000 --- a/lib/anubis/server/registry/adapter.ex +++ /dev/null @@ -1,164 +0,0 @@ -defmodule Anubis.Server.Registry.Adapter do - @moduledoc """ - Behaviour for registry adapters in MCP servers. - - This module defines the interface that registry implementations must follow - to be pluggable into the Anubis MCP server architecture. It allows users - to provide custom registry implementations (e.g., using Horde for cluster-wide - distribution) while maintaining compatibility with the existing API. - - ## Implementing a Custom Registry - - To implement a custom registry adapter, create a module that implements - all the callbacks defined in this behaviour: - - defmodule MyApp.HordeRegistry do - @behaviour Anubis.Server.Registry.Adapter - - def child_spec(opts) do - %{ - id: __MODULE__, - start: {Horde.Registry, :start_link, [ - [ - name: __MODULE__, - keys: :unique, - members: :auto - ] ++ opts - ]} - } - end - - def server(module) do - {:via, Horde.Registry, {__MODULE__, {:server, module}}} - end - - # ... implement other callbacks - end - - ## Using a Custom Registry - - You can configure a custom registry at multiple levels: - - Anubis.Server.start_link(MyServer, :ok, transport: :stdio, registry: MyApp.HordeRegistry) - - ## Default Implementation - - The default implementation uses Elixir's built-in Registry module. - """ - - @doc """ - Returns a child specification for the registry. - - This is used when starting the registry as part of a supervision tree. - The implementation should return a valid child specification map or tuple. - """ - @callback child_spec(opts :: keyword()) :: Supervisor.child_spec() - - @doc """ - Returns a name for a server process. - - The returned value must be a valid GenServer name that can be passed - to `GenServer.start_link/3` and similar functions. - - ## Parameters - - * `module` - The module implementing the server - """ - @callback server(module :: module()) :: GenServer.name() - - @doc """ - Returns a name for a `Task.Supervisor` process. - - The returned value must be a valid GenServer name that can be passed - to `GenServer.start_link/3` and similar functions. - - ## Parameters - - * `module` - The module implementing the server - """ - @callback task_supervisor(module :: module()) :: GenServer.name() - - @doc """ - Returns a name for a server session process. - - ## Parameters - - * `server_module` - The module implementing the server - * `session_id` - The unique session identifier - """ - @callback server_session(server_module :: module(), session_id :: String.t()) :: - GenServer.name() - - @doc """ - Returns a name for a transport process. - - ## Parameters - - * `server_module` - The module implementing the server - * `transport_type` - The type of transport (e.g., :stdio, :sse, :websocket) - """ - @callback transport(server_module :: module(), transport_type :: atom()) :: - GenServer.name() - - @doc """ - Returns a name for a supervisor process. - - ## Parameters - - * `kind` - The kind of supervisor (e.g., :supervisor, :session_supervisor) - * `server_module` - The module implementing the server - """ - @callback supervisor(kind :: atom(), server_module :: module()) :: GenServer.name() - - @doc """ - Gets the PID of a registered server. - - Returns the PID if the server is registered, nil otherwise. - - ## Parameters - - * `server_module` - The module implementing the server - """ - @callback whereis_server(server_module :: module()) :: pid() | nil - - @doc """ - Gets the PID of a server session process. - - Returns the PID if the session is registered, nil otherwise. - - ## Parameters - - * `server_module` - The module implementing the server - * `session_id` - The unique session identifier - """ - @callback whereis_server_session( - server_module :: module(), - session_id :: String.t() - ) :: pid() | nil - - @doc """ - Gets the PID of a transport process. - - Returns the PID if the transport is registered, nil otherwise. - - ## Parameters - - * `server_module` - The module implementing the server - * `transport_type` - The type of transport - """ - @callback whereis_transport(server_module :: module(), transport_type :: atom()) :: - pid() | nil - - @doc """ - Gets the PID of a supervisor process. - - Returns the PID if the supervisor is registered, nil otherwise. - - ## Parameters - - * `kind` - The kind of supervisor - * `server_module` - The module implementing the server - """ - @callback whereis_supervisor(kind :: atom(), server_module :: module()) :: - pid() | nil -end diff --git a/lib/anubis/server/registry/local.ex b/lib/anubis/server/registry/local.ex new file mode 100644 index 00000000..de9bb40b --- /dev/null +++ b/lib/anubis/server/registry/local.ex @@ -0,0 +1,100 @@ +defmodule Anubis.Server.Registry.Local do + @moduledoc """ + ETS-based session registry for HTTP transports. + + Uses a named ETS table with `read_concurrency: true` for fast lookups. + Monitors registered processes for automatic cleanup on crash/shutdown. + """ + + @behaviour Anubis.Server.Registry + + use GenServer + + @impl Anubis.Server.Registry + def child_spec(opts) do + name = Keyword.fetch!(opts, :name) + + %{ + id: {__MODULE__, name}, + start: {__MODULE__, :start_link, [opts]}, + type: :worker, + restart: :permanent + } + end + + def start_link(opts) do + name = Keyword.fetch!(opts, :name) + GenServer.start_link(__MODULE__, opts, name: name) + end + + @impl Anubis.Server.Registry + def register_session(name, session_id, pid) do + GenServer.call(name, {:register, session_id, pid}) + end + + @impl Anubis.Server.Registry + def lookup_session(name, session_id) do + table = table_name(name) + + case :ets.lookup(table, session_id) do + [{^session_id, pid}] when is_pid(pid) -> + if Process.alive?(pid), do: {:ok, pid}, else: {:error, :not_found} + + [] -> + {:error, :not_found} + end + rescue + ArgumentError -> {:error, :not_found} + end + + @impl Anubis.Server.Registry + def unregister_session(name, session_id) do + GenServer.call(name, {:unregister, session_id}) + end + + # GenServer callbacks + + @impl GenServer + def init(opts) do + name = Keyword.fetch!(opts, :name) + table = table_name(name) + ^table = :ets.new(table, [:named_table, :public, :set, read_concurrency: true]) + + {:ok, %{table: table, monitors: %{}}} + end + + @impl GenServer + def handle_call({:register, session_id, pid}, _from, state) do + :ets.insert(state.table, {session_id, pid}) + ref = Process.monitor(pid) + monitors = Map.put(state.monitors, ref, session_id) + {:reply, :ok, %{state | monitors: monitors}} + end + + def handle_call({:unregister, session_id}, _from, state) do + :ets.delete(state.table, session_id) + + monitors = + state.monitors + |> Enum.reject(fn {_ref, sid} -> sid == session_id end) + |> Map.new() + + {:reply, :ok, %{state | monitors: monitors}} + end + + @impl GenServer + def handle_info({:DOWN, ref, :process, _pid, _reason}, state) do + case Map.pop(state.monitors, ref) do + {nil, monitors} -> + {:noreply, %{state | monitors: monitors}} + + {session_id, monitors} -> + :ets.delete(state.table, session_id) + {:noreply, %{state | monitors: monitors}} + end + end + + def handle_info(_msg, state), do: {:noreply, state} + + defp table_name(name) when is_atom(name), do: :"#{name}.ets" +end diff --git a/lib/anubis/server/registry/none.ex b/lib/anubis/server/registry/none.ex new file mode 100644 index 00000000..a6bbfb8c --- /dev/null +++ b/lib/anubis/server/registry/none.ex @@ -0,0 +1,21 @@ +defmodule Anubis.Server.Registry.None do + @moduledoc """ + No-op registry for STDIO transport. + + STDIO has exactly one session, looked up by atom name. No registry needed. + """ + + @behaviour Anubis.Server.Registry + + @impl Anubis.Server.Registry + def child_spec(_opts), do: :ignore + + @impl Anubis.Server.Registry + def register_session(_name, _session_id, _pid), do: :ok + + @impl Anubis.Server.Registry + def lookup_session(_name, _session_id), do: {:error, :not_found} + + @impl Anubis.Server.Registry + def unregister_session(_name, _session_id), do: :ok +end diff --git a/lib/anubis/server/session.ex b/lib/anubis/server/session.ex index 437319fe..d34429e3 100644 --- a/lib/anubis/server/session.ex +++ b/lib/anubis/server/session.ex @@ -1,271 +1,972 @@ defmodule Anubis.Server.Session do - @moduledoc false + @moduledoc """ + Per-client MCP session process. - use Agent, restart: :transient + Each Session is a GenServer that manages the lifecycle of a single MCP client + connection. It handles protocol initialization, request/notification dispatch, + server-initiated requests (sampling, roots), and session persistence. + + Sessions are created by the transport layer (STDIO creates one at startup, + HTTP transports create them dynamically via `Anubis.Server.Supervisor`). + """ + + use GenServer use Anubis.Logging import Peri - @type t :: %__MODULE__{ + alias Anubis.MCP.Error + alias Anubis.MCP.ID + alias Anubis.MCP.Message + alias Anubis.Server + alias Anubis.Server.Context + alias Anubis.Server.Frame + alias Anubis.Telemetry + + require Message + require Server + + @default_session_idle_timeout to_timeout(minute: 30) + + @type t :: %{ + session_id: String.t(), + server_module: module(), protocol_version: String.t() | nil, protocol_module: module() | nil, initialized: boolean(), - name: GenServer.name() | nil, client_info: map() | nil, client_capabilities: map() | nil, - log_level: String.t(), - id: String.t() | nil, + log_level: String.t() | nil, + frame: Frame.t(), + server_info: map(), + capabilities: map(), + supported_versions: list(String.t()), + transport: %{layer: module(), name: GenServer.name()}, + registry: module(), + session_idle_timeout: pos_integer(), + expiry_timer: reference() | nil, pending_requests: %{ String.t() => %{started_at: integer(), method: String.t()} - } + }, + server_requests: %{ + String.t() => %{ + method: String.t(), + timer_ref: reference() + } + }, + timeout: pos_integer(), + task_supervisor: GenServer.name() } - defstruct [ - :id, - :protocol_version, - :protocol_module, - :log_level, - :name, - initialized: false, - client_info: nil, - client_capabilities: nil, - pending_requests: %{} - ] - - defschema :state_t, %{ - protocol_version: :string, - protocol_module: :atom, - initialized: {:required, :boolean}, - name: {:custom, &Anubis.genserver_name/1}, - client_info: :map, - client_capabilities: :map, - log_level: :string, - id: :string, - pending_requests: {:map, :string, %{started_at: :integer, method: :string}} - } + defschema(:parse_options, [ + {:session_id, {:required, :string}}, + {:server_module, {:required, :atom}}, + {:name, {:required, {:custom, &Anubis.genserver_name/1}}}, + {:transport, {:required, {:custom, &Anubis.server_transport/1}}}, + {:registry, {:atom, {:default, Anubis.Server.Registry}}}, + {:session_idle_timeout, {{:integer, {:gte, 1}}, {:default, @default_session_idle_timeout}}}, + {:timeout, {:integer, {:default, to_timeout(second: 30)}}}, + {:task_supervisor, {:required, {:custom, &Anubis.genserver_name/1}}} + ]) @doc """ - Starts a new session agent with initial state. - - If a session store is configured and the session exists in storage, - it will be restored. Otherwise, a new session is created. + Starts a Session process linked to the current process. + + ## Options + + * `:session_id` — unique session identifier (required) + * `:server_module` — the MCP server module implementing `Anubis.Server` (required) + * `:name` — GenServer registration name (required) + * `:transport` — transport configuration `[layer: module, name: name]` (required) + * `:task_supervisor` — name of the `Task.Supervisor` for async work (required) + * `:registry` — session registry module (default: `Anubis.Server.Registry`) + * `:session_idle_timeout` — idle timeout in ms before session expires (default: 30 min) + * `:timeout` — request timeout in ms (default: 30s) """ - @spec start_link(keyword()) :: Agent.on_start() - def start_link(opts \\ []) do - session_id = Keyword.fetch!(opts, :session_id) + @spec start_link(keyword()) :: GenServer.on_start() + def start_link(opts) do + opts = parse_options!(opts) name = Keyword.fetch!(opts, :name) - server_module = Keyword.get(opts, :server_module) - - initial_state = - case maybe_restore_session(session_id, name, server_module) do - {:ok, state} -> - Logging.log(:info, "Restored session #{inspect(session_id)} from store", - initialized: state.initialized, - protocol_version: state.protocol_version - ) - state + GenServer.start_link(__MODULE__, Map.new(opts), name: name) + end + + # Lifecycle + + @impl GenServer + def init(opts) do + module = opts.server_module + server_info = module.server_info() + capabilities = module.server_capabilities() + protocol_versions = module.supported_protocol_versions() + + state = %{ + session_id: opts.session_id, + server_module: module, + protocol_version: nil, + protocol_module: nil, + initialized: false, + client_info: nil, + client_capabilities: nil, + log_level: nil, + frame: Frame.new(), + server_info: server_info, + capabilities: capabilities, + supported_versions: protocol_versions, + transport: Map.new(opts.transport), + registry: opts.registry, + session_idle_timeout: opts.session_idle_timeout, + expiry_timer: nil, + pending_requests: %{}, + server_requests: %{}, + timeout: opts.timeout, + task_supervisor: opts.task_supervisor + } + + state = schedule_session_expiry(state) + + Logging.server_event("session_starting", %{ + session_id: opts.session_id, + module: module, + server_info: server_info + }) + + Telemetry.execute( + Telemetry.event_server_init(), + %{system_time: System.system_time()}, + %{ + module: module, + server_info: server_info, + capabilities: capabilities, + session_id: opts.session_id + } + ) + + {:ok, state, :hibernate} + end + + # Request/Response handling - {:error, _reason} -> - new(id: session_id, name: name) + @impl GenServer + def handle_call({:mcp_request, decoded, transport_context}, _from, state) when is_map(decoded) do + state = merge_transport_assigns(state, transport_context) + state = reset_session_expiry(state) + + handle_single_request(decoded, transport_context, state) + end + + def handle_call(request, from, %{server_module: module} = state) do + if Anubis.exported?(module, :handle_call, 3) do + frame = prepare_frame(state) + + case module.handle_call(request, from, frame) do + {:reply, reply, frame} -> + {:reply, reply, %{state | frame: frame}} + + {:reply, reply, frame, cont} -> + {:reply, reply, %{state | frame: frame}, cont} + + {:noreply, frame} -> + {:noreply, %{state | frame: frame}} + + {:noreply, frame, cont} -> + {:noreply, %{state | frame: frame}, cont} + + {:stop, reason, reply, frame} -> + {:stop, reason, reply, %{state | frame: frame}} + + {:stop, reason, frame} -> + {:stop, reason, %{state | frame: frame}} end + else + {:reply, {:error, :not_implemented}, state} + end + end + + # Notification dispatch - Agent.start_link(fn -> initial_state end, name: name) + @impl GenServer + def handle_cast({:mcp_notification, decoded, transport_context}, state) when is_map(decoded) do + state = merge_transport_assigns(state, transport_context) + state = reset_session_expiry(state) + + if Message.is_initialize_lifecycle(decoded) or state.initialized do + handle_notification(decoded, transport_context, state) + else + Logging.server_event("session_not_initialized_check", %{ + session_id: state.session_id, + initialized: state.initialized, + method: decoded["method"] + }) + + {:noreply, state} + end end - @doc """ - Creates a new server state with the given options. - """ - @spec new(Enumerable.t()) :: t() - def new(opts), do: struct(__MODULE__, opts) + # Server-initiated request responses (sampling/roots) - @doc """ - Guard to check if a session has been initialized. - """ - defguard is_initialized(session) when session.initialized + def handle_cast({:mcp_response, decoded, _context}, state) when is_map(decoded) do + cond do + Message.is_response(decoded) and server_request?(decoded["id"], state) -> + handle_server_request_response(decoded, state) - @doc """ - Retrieves the current state of a session. - """ - @spec get(GenServer.name()) :: t - def get(session) do - Agent.get(session, & &1) + Message.is_error(decoded) and server_request?(decoded["id"], state) -> + handle_server_request_error(decoded, state) + + true -> + Logging.server_event( + "unexpected_response", + %{message: decoded}, + level: :warning + ) + + {:noreply, state} + end end - @doc """ - Updates state after successful initialization handshake. + def handle_cast(request, %{server_module: module} = state) do + if Anubis.exported?(module, :handle_cast, 2) do + frame = prepare_frame(state) - This function: - 1. Sets the negotiated protocol version - 2. Stores client information and capabilities - 3. Persists the session if a store is configured + case module.handle_cast(request, frame) do + {:noreply, frame} -> {:noreply, %{state | frame: frame}} + {:noreply, frame, cont} -> {:noreply, %{state | frame: frame}, cont} + {:stop, reason, frame} -> {:stop, reason, %{state | frame: frame}} + end + else + {:noreply, state} + end + end - Note: Call `mark_initialized/1` separately to set the initialized flag. - """ - @spec update_from_initialization(GenServer.name(), String.t(), map, map, keyword()) :: :ok - def update_from_initialization(session, negotiated_version, client_info, capabilities, opts \\ []) do - protocol_module = Keyword.get(opts, :protocol_module) - - Agent.update(session, fn state -> - new_state = %{ - state - | protocol_version: negotiated_version, - protocol_module: protocol_module || state.protocol_module, - client_info: client_info, - client_capabilities: capabilities - } + # Handle info messages - maybe_persist_session(new_state) - new_state - end) + @impl GenServer + def handle_info({:send_notification, method, params}, state) do + with {:ok, notification} <- encode_notification(method, params), + :ok <- send_to_transport(state.transport, notification, timeout: state.timeout) do + {:noreply, state} + else + {:error, err} -> + Logging.server_event("failed_send_notification", %{method: method, error: err}, level: :error) + + {:noreply, state} + end end - @doc """ - Marks the session as initialized. - """ - @spec mark_initialized(GenServer.name()) :: :ok - def mark_initialized(session) do - Agent.update(session, fn state -> - new_state = %{state | initialized: true} - maybe_persist_session(new_state) - new_state - end) + def handle_info(:session_expired, state) do + Logging.server_event("session_expired", %{session_id: state.session_id}) + {:stop, {:shutdown, :session_expired}, state} end - @doc """ - Updates the log level. - """ - @spec set_log_level(GenServer.name(), String.t()) :: :ok - def set_log_level(session, level) do - Agent.update(session, fn state -> %{state | log_level: level} end) + def handle_info({:send_sampling_request, params, timeout}, state) do + request_id = ID.generate_request_id() + handle_sampling_request_send(request_id, params, timeout, state) end - @doc """ - Tracks a new pending request in the session. - """ - @spec track_request(GenServer.name(), String.t(), String.t()) :: :ok - def track_request(session, request_id, method) do - Agent.update(session, fn state -> - request_info = %{ - started_at: System.system_time(:millisecond), - method: method - } + def handle_info({:sampling_request_timeout, request_id}, state) do + handle_sampling_timeout(request_id, state) + end - %{ - state - | pending_requests: Map.put(state.pending_requests, request_id, request_info) - } - end) + def handle_info({:send_roots_request, timeout}, state) do + request_id = ID.generate_request_id() + handle_roots_request_send(request_id, timeout, state) end - @doc """ - Removes a completed request from tracking. - """ - @spec complete_request(GenServer.name(), String.t()) :: map() | nil - def complete_request(session, request_id) do - Agent.get_and_update(session, fn state -> - {request_info, pending_requests} = Map.pop(state.pending_requests, request_id) - {request_info, %{state | pending_requests: pending_requests}} - end) + def handle_info({:roots_request_timeout, request_id}, state) do + handle_roots_timeout(request_id, state) end - @doc """ - Checks if a request is currently pending. - """ - @spec has_pending_request?(GenServer.name(), String.t()) :: boolean() - def has_pending_request?(session, request_id) do - Agent.get(session, fn state -> - Map.has_key?(state.pending_requests, request_id) + def handle_info(event, %{server_module: module} = state) do + if Anubis.exported?(module, :handle_info, 2) do + frame = prepare_frame(state) + + case module.handle_info(event, frame) do + {:noreply, frame} -> {:noreply, %{state | frame: frame}} + {:noreply, frame, cont} -> {:noreply, %{state | frame: frame}, cont} + {:stop, reason, frame} -> {:stop, reason, %{state | frame: frame}} + end + else + {:noreply, state} + end + end + + @impl GenServer + def terminate(reason, %{server_module: module, server_info: server_info} = state) do + cancel_session_expiry(state) + + Logging.server_event("session_terminating", %{ + session_id: state.session_id, + reason: reason, + server_info: server_info + }) + + Telemetry.execute( + Telemetry.event_server_terminate(), + %{system_time: System.system_time()}, + %{reason: reason, server_info: server_info, session_id: state.session_id} + ) + + if Anubis.exported?(module, :terminate, 2) do + frame = prepare_frame(state) + module.terminate(reason, frame) + else + :ok + end + end + + @impl GenServer + def format_status(status) do + Map.new(status, fn + {:state, state} -> + {:state, format_state(state)} + + {:message, {:mcp_request, decoded, _ctx}} -> + {:message, {:mcp_request, decoded}} + + {:message, {:mcp_notification, decoded, _ctx}} -> + {:message, {:mcp_notification, decoded}} + + {:message, {:mcp_response, decoded, _ctx}} -> + {:message, {:mcp_response, decoded}} + + other -> + other end) end - @doc """ - Gets all pending requests for a session. - """ - @spec get_pending_requests(GenServer.name()) :: map() - def get_pending_requests(session) do - Agent.get(session, & &1.pending_requests) + # Request handling + + defguardp is_server_initialized(decoded, state) + when Message.is_initialize_lifecycle(decoded) or + state.initialized == true + + defp handle_single_request(decoded, transport_context, state) do + cond do + Message.is_response(decoded) and server_request?(decoded["id"], state) -> + handle_server_request_response(decoded, state) + + Message.is_error(decoded) and server_request?(decoded["id"], state) -> + handle_server_request_error(decoded, state) + + Message.is_ping(decoded) -> + handle_server_ping(decoded, state) + + not is_server_initialized(decoded, state) -> + handle_server_not_initialized(state) + + Message.is_request(decoded) -> + handle_request(decoded, transport_context, state) + + true -> + handle_invalid_request(state) + end end - # Private persistence functions + defp handle_server_ping(%{"id" => request_id}, state) do + {:reply, {:ok, encode_reply(Message.build_response(%{}, request_id))}, state} + end - defp maybe_restore_session(session_id, name, server_module) do - if store = Anubis.get_session_store_adapter() do - Logging.log(:debug, "Attempting to restore session from store. session_id: #{inspect(session_id)}", []) - - case store.load(session_id, server: server_module) do - {:ok, state_map} -> - Logging.log(:debug, "Successfully loaded session #{inspect(session_id)} from store", []) - {:ok, state} = state_t(state_map) - state = struct(__MODULE__, state) - {:ok, %{state | name: name}} - - {:error, :not_found} = error -> - Logging.log(:debug, "Session #{inspect(session_id)} not found in store, creating new session", []) - error - - error -> - Logging.log(:debug, "Failed to load session #{inspect(session_id)} from store", error: error) - error + defp handle_server_not_initialized(state) do + error = Error.protocol(:invalid_request, %{message: "Server not initialized"}) + + Logging.server_event( + "request_error", + %{error: error, reason: "not_initialized"}, + level: :warning + ) + + {:reply, {:ok, encode_reply(Error.build_json_rpc(error))}, state} + end + + defp handle_invalid_request(state) do + error = + Error.protocol(:invalid_request, %{ + message: "Expected request but got different message type" + }) + + {:reply, {:error, error}, state} + end + + # Initialize handling + + defp handle_request(%{"params" => params} = request, _transport_context, state) when Message.is_initialize(request) do + %{ + "clientInfo" => client_info, + "capabilities" => client_capabilities, + "protocolVersion" => requested_version + } = params + + {:ok, protocol_version, protocol_module} = + Anubis.Protocol.Registry.negotiate(requested_version, state.supported_versions) + + state = %{ + state + | protocol_version: protocol_version, + protocol_module: protocol_module, + client_info: client_info, + client_capabilities: client_capabilities + } + + maybe_persist_session(state) + + result = %{ + "protocolVersion" => protocol_version, + "serverInfo" => state.server_info, + "capabilities" => state.capabilities + } + + Logging.server_event("initializing", %{ + client_info: client_info, + client_capabilities: client_capabilities, + protocol_version: protocol_version, + session_id: state.session_id + }) + + Telemetry.execute( + Telemetry.event_server_response(), + %{system_time: System.system_time()}, + %{method: "initialize", status: :success} + ) + + {:reply, {:ok, encode_reply(Message.build_response(result, request["id"]))}, state} + end + + defp handle_request(%{"id" => request_id, "method" => "logging/setLevel"} = request, _transport_context, state) + when Server.is_supported_capability(state.capabilities, "logging") do + level = request["params"]["level"] + state = %{state | log_level: level} + {:reply, {:ok, encode_reply(Message.build_response(%{}, request_id))}, state} + end + + defp handle_request(%{"id" => request_id, "method" => method} = request, transport_context, state) do + Logging.server_event("handling_request", %{ + id: request_id, + method: method, + session_id: state.session_id + }) + + state = track_request(state, request_id, method) + + Telemetry.execute( + Telemetry.event_server_request(), + %{system_time: System.system_time()}, + %{id: request_id, method: method} + ) + + frame = prepare_frame(state, transport_context) + server_request(request, %{state | frame: frame}) + end + + # Notification handling + + defp handle_notification( + %{"method" => "notifications/initialized"}, + _transport_context, + %{server_module: module} = state + ) do + Logging.server_event("client_initialized", %{session_id: state.session_id}) + + state = %{state | initialized: true} + + maybe_persist_session(state) + + Logging.server_event("session_marked_initialized", %{ + session_id: state.session_id, + initialized: true + }) + + frame = prepare_frame(state) + + {:ok, frame} = + if Anubis.exported?(module, :init, 2), + do: module.init(state.client_info, frame), + else: {:ok, frame} + + {:noreply, %{state | frame: frame}} + end + + defp handle_notification(%{"method" => "notifications/cancelled"} = notification, _transport_context, state) do + params = notification["params"] || %{} + request_id = params["requestId"] + reason = Map.get(params, "reason", "cancelled") + + case Map.get(state.pending_requests, request_id) do + nil -> + Logging.server_event("cancellation_for_unknown_request", %{ + session_id: state.session_id, + request_id: request_id, + reason: reason + }) + + {:noreply, state} + + request_info -> + state = complete_request(state, request_id) + + Logging.server_event("request_cancelled", %{ + session_id: state.session_id, + request_id: request_id, + reason: reason, + method: request_info[:method], + duration_ms: System.system_time(:millisecond) - request_info[:started_at] + }) + + Telemetry.execute( + Telemetry.event_server_notification(), + %{system_time: System.system_time()}, + %{ + method: "cancelled", + session_id: state.session_id, + request_id: request_id + } + ) + + {:noreply, state} + end + end + + defp handle_notification(notification, _transport_context, state) do + method = notification["method"] + + Logging.server_event("handling_notification", %{method: method}) + + Telemetry.execute( + Telemetry.event_server_notification(), + %{system_time: System.system_time()}, + %{method: method} + ) + + frame = prepare_frame(state) + server_notification(notification, %{state | frame: frame}) + end + + # Server request/notification dispatch + + defp server_request(%{"id" => request_id, "method" => method} = request, %{server_module: module} = state) do + case module.handle_request(request, state.frame) do + {:reply, response, %Frame{} = frame} -> + Telemetry.execute( + Telemetry.event_server_response(), + %{system_time: System.system_time()}, + %{id: request_id, method: method, status: :success} + ) + + state = complete_request(%{state | frame: frame}, request_id) + + {:reply, {:ok, encode_reply(Message.build_response(response, request_id))}, state} + + {:noreply, %Frame{} = frame} -> + Telemetry.execute( + Telemetry.event_server_response(), + %{system_time: System.system_time()}, + %{id: request_id, method: method, status: :noreply} + ) + + state = complete_request(%{state | frame: frame}, request_id) + {:reply, {:ok, nil}, state} + + {:error, %Error{} = error, %Frame{} = frame} -> + Logging.server_event( + "request_error", + %{id: request_id, method: method, error: error}, + level: :warning + ) + + Telemetry.execute( + Telemetry.event_server_error(), + %{system_time: System.system_time()}, + %{id: request_id, method: method, error: error} + ) + + state = complete_request(%{state | frame: frame}, request_id) + + {:reply, {:ok, encode_reply(Error.build_json_rpc(error, request_id))}, state} + end + end + + defp server_notification(%{"method" => method} = notification, %{server_module: module} = state) do + if Anubis.exported?(module, :handle_notification, 2) do + case module.handle_notification(notification, state.frame) do + {:noreply, %Frame{} = frame} -> + {:noreply, %{state | frame: frame}} + + {:error, _error, %Frame{} = frame} -> + Logging.server_event( + "notification_handler_error", + %{method: method}, + level: :warning + ) + + {:noreply, %{state | frame: frame}} end else - Logging.log(:debug, "No session store configured, creating new session", session_id: session_id) + {:noreply, state} + end + end + + # Request tracking + + defp track_request(state, request_id, method) do + request_info = %{ + started_at: System.system_time(:millisecond), + method: method + } + + %{state | pending_requests: Map.put(state.pending_requests, request_id, request_info)} + end + + defp complete_request(state, request_id) do + %{state | pending_requests: Map.delete(state.pending_requests, request_id)} + end + + # Frame management + + defp prepare_frame(state, transport_context \\ nil) do + headers = + case transport_context do + %{req_headers: req_headers} -> normalize_headers(req_headers) + _ -> %{} + end + + remote_ip = + case transport_context do + %{remote_ip: ip} -> ip + _ -> nil + end + + context = %Context{ + session_id: state.session_id, + client_info: state.client_info, + headers: headers, + remote_ip: remote_ip + } + + %{state.frame | context: context} + end + + defp merge_transport_assigns(state, %{assigns: assigns}) when is_map(assigns) do + original_context = state.frame.context + frame = Frame.assign(state.frame, assigns) + frame = %{frame | context: original_context} + %{state | frame: frame} + end + + defp merge_transport_assigns(state, _context), do: state + + defp normalize_headers(req_headers) when is_list(req_headers) do + Map.new(req_headers, fn {k, v} -> {String.downcase(k), v} end) + end + + defp normalize_headers(_), do: %{} - {:error, :no_store} + # Session expiry management + + defp schedule_session_expiry(%{session_idle_timeout: timeout} = state) do + timer = Process.send_after(self(), :session_expired, timeout) + %{state | expiry_timer: timer} + end + + defp reset_session_expiry(state) do + cancel_session_expiry(state) + schedule_session_expiry(state) + end + + defp cancel_session_expiry(%{expiry_timer: timer} = state) do + if timer, do: Process.cancel_timer(timer) + %{state | expiry_timer: nil} + end + + # Reply encoding + + defp encode_reply(message) when is_map(message) do + JSON.encode!(message) + end + + # Transport helpers + + defp encode_notification(method, params) do + notification = Message.build_notification(method, params) + Logging.message("outgoing", "notification", nil, notification) + Message.encode_notification(notification) + end + + defp send_to_transport(nil, _data, _opts) do + {:error, Error.transport(:no_transport, %{message: "No transport configured"})} + end + + defp send_to_transport(%{layer: layer, name: name}, data, opts) do + with {:error, reason} <- layer.send_message(name, data, opts) do + {:error, Error.transport(:send_failure, %{original_reason: reason})} end end - defp maybe_persist_session(%__MODULE__{} = state) do - if store = Anubis.get_session_store_adapter() do - Logging.log(:debug, "Persisting session #{inspect(state.id)} to store", []) + # Sampling request helpers + + defp handle_sampling_request_send(request_id, params, timeout, state) do + timer_ref = + Process.send_after(self(), {:sampling_request_timeout, request_id}, timeout) + + request_info = %{ + method: "sampling/createMessage", + session_id: state.session_id, + timer_ref: timer_ref + } + + state = put_in(state.server_requests[request_id], request_info) + + with :ok <- validate_client_capability(state, "sampling"), + {:ok, request_data} <- + encode_request("sampling/createMessage", params, request_id), + :ok <- send_to_transport(state.transport, request_data, timeout: state.timeout) do + Logging.server_event("sent_sampling_request", %{request_id: request_id}) + {:noreply, state} + else + {:error, error} -> + Process.cancel_timer(timer_ref) + + state = %{ + state + | server_requests: Map.delete(state.server_requests, request_id) + } + + Logging.server_event( + "failed_send_sampling_request", + %{request_id: request_id, error: error}, + level: :error + ) + + {:noreply, state} + end + end + + defp validate_client_capability(state, capability) do + if Map.has_key?(state.client_capabilities || %{}, capability) do + :ok + else + {:error, "Client does not support #{capability} capability"} + end + end + + defp handle_sampling_timeout(request_id, state) do + case Map.pop(state.server_requests, request_id) do + {nil, _} -> + {:noreply, state} - # Convert struct to mapand remove runtime fields - state_map = - state - |> Map.from_struct() - # Don't persist process names - |> Map.delete(:name) + {_request_info, updated_requests} -> + Logging.server_event("sampling_request_timeout", %{request_id: request_id}, level: :warning) - case store.save(state.id, state_map, []) do + {:noreply, %{state | server_requests: updated_requests}} + end + end + + defp encode_request(method, params, request_id) do + request = %{ + "method" => method, + "params" => params + } + + Logging.message("outgoing", "request", request_id, request) + Message.encode_request(request, request_id) + end + + defp server_request?(request_id, %{server_requests: requests}) when is_binary(request_id) do + Map.has_key?(requests, request_id) + end + + defp server_request?(_, _), do: false + + defp handle_server_request_response(%{"id" => request_id, "result" => result}, state) do + {request_info, updated_requests} = Map.pop(state.server_requests, request_id) + Process.cancel_timer(request_info.timer_ref) + + state = %{state | server_requests: updated_requests} + + case request_info.method do + "sampling/createMessage" -> + handle_sampling(result, request_id, state) + + "roots/list" -> + handle_roots(result["roots"] || [], request_id, state) + + _ -> + {:noreply, state} + end + end + + defp handle_server_request_error(%{"id" => request_id, "error" => error}, state) do + {request_info, updated_requests} = Map.pop(state.server_requests, request_id) + Process.cancel_timer(request_info.timer_ref) + + state = %{state | server_requests: updated_requests} + + Logging.server_event( + "server_request_error", + %{ + request_id: request_id, + method: request_info.method, + error: error + }, + level: :error + ) + + {:noreply, state} + end + + defp handle_sampling(result, request_id, %{server_module: module} = state) do + if Anubis.exported?(module, :handle_sampling, 3) do + frame = prepare_frame(state) + + case module.handle_sampling(result, request_id, frame) do + {:noreply, new_frame} -> + {:noreply, %{state | frame: new_frame}} + + {:stop, reason, new_frame} -> + {:stop, reason, %{state | frame: new_frame}} + end + else + {:noreply, state} + end + end + + # Roots request helpers + + defp handle_roots_request_send(request_id, timeout, state) do + timer_ref = + Process.send_after(self(), {:roots_request_timeout, request_id}, timeout) + + request_info = %{ + id: request_id, + method: "roots/list", + session_id: state.session_id, + timer_ref: timer_ref + } + + state = put_in(state.server_requests[request_id], request_info) + + with :ok <- validate_client_capability(state, "roots"), + {:ok, request_data} <- encode_request("roots/list", %{}, request_id), + :ok <- send_to_transport(state.transport, request_data, timeout: state.timeout) do + Logging.server_event("sent_roots_request", %{request_id: request_id}) + {:noreply, state} + else + {:error, error} -> + Process.cancel_timer(timer_ref) + + state = %{ + state + | server_requests: Map.delete(state.server_requests, request_id) + } + + Logging.server_event( + "failed_send_roots_request", + %{request_id: request_id, error: error}, + level: :error + ) + + {:noreply, state} + end + end + + defp handle_roots_timeout(request_id, state) when is_binary(request_id) do + state.server_requests + |> Map.pop(request_id) + |> handle_roots_timeout(state) + end + + defp handle_roots_timeout({nil, _}, state), do: {:noreply, state} + + defp handle_roots_timeout({%{id: request_id}, requests}, state) do + with {:ok, notification} <- + encode_notification("notifications/cancelled", %{ + "requestId" => request_id, + "reason" => "timeout" + }), + :ok <- send_to_transport(state.transport, notification, timeout: state.timeout) do + Logging.server_event( + "roots_request_timeout_cancelled", + %{request_id: request_id} + ) + end + + Logging.server_event("roots_request_timeout", %{request_id: request_id}, level: :warning) + + {:noreply, %{state | server_requests: requests}} + end + + defp handle_roots(roots, request_id, %{server_module: module} = state) do + if Anubis.exported?(module, :handle_roots, 3) do + frame = prepare_frame(state) + + case module.handle_roots(roots, request_id, frame) do + {:noreply, new_frame} -> + {:noreply, %{state | frame: new_frame}} + + {:stop, reason, new_frame} -> + {:stop, reason, %{state | frame: new_frame}} + end + else + {:noreply, state} + end + end + + # Session persistence + + defp maybe_persist_session(%{session_id: session_id} = state) do + if store = Anubis.get_session_store_adapter() do + Logging.log(:debug, "Persisting session #{inspect(session_id)} to store", []) + + state_map = %{ + protocol_version: state.protocol_version, + protocol_module: state.protocol_module, + initialized: state.initialized, + client_info: state.client_info, + client_capabilities: state.client_capabilities, + log_level: state.log_level, + id: session_id, + pending_requests: state.pending_requests, + frame: Frame.to_saved(state.frame) + } + + case store.save(session_id, state_map, []) do :ok -> - Logging.log(:debug, "Successfully persisted session #{inspect(state.id)} to store", []) + Logging.log(:debug, "Successfully persisted session #{inspect(session_id)}", []) {:error, reason} -> Logging.log( :warning, - "Failed to persist session #{inspect(state.id)} to store", - session_id: state.id, + "Failed to persist session #{inspect(session_id)}", + session_id: session_id, error: reason ) :ok end - else - Logging.log(:debug, "No session store configured, skipping persistence", []) end end -end - -defimpl Inspect, for: Anubis.Server.Session do - import Inspect.Algebra - def inspect(session, opts) do - info = [ - id: session.id, - initialized: session.initialized, - pending_requests: map_size(session.pending_requests) - ] - - info = - if session.protocol_version, - do: [{:protocol_version, session.protocol_version} | info], - else: info - - info = - if session.client_info, - do: [{:client_info, session.client_info["name"] || "unknown"} | info], - else: info + # Format helpers + + defp format_state(state) do + pending = format_pending_requests(state.server_requests) + + state + |> Map.take([ + :session_id, + :server_module, + :initialized, + :protocol_version, + :capabilities, + :frame + ]) + |> Map.merge(%{ + transport: state.transport[:layer], + pending_server_requests: pending + }) + end - concat(["#Session<", to_doc(info, opts), ">"]) + defp format_pending_requests(requests) do + Enum.map(requests, fn {id, req} -> + %{id: id, method: req[:method]} + end) end end diff --git a/lib/anubis/server/session/supervisor.ex b/lib/anubis/server/session/supervisor.ex deleted file mode 100644 index 50042e3e..00000000 --- a/lib/anubis/server/session/supervisor.ex +++ /dev/null @@ -1,127 +0,0 @@ -defmodule Anubis.Server.Session.Supervisor do - @moduledoc false - - use DynamicSupervisor - use Anubis.Logging - - alias Anubis.Server.Session - - @kind :session_supervisor - - @doc """ - Starts the session supervisor. - - ## Parameters - * `server` - The server module atom - - ## Returns - * `{:ok, pid}` - Supervisor started successfully - * `{:error, reason}` - Failed to start supervisor - - ## Examples - - {:ok, _pid} = Session.Supervisor.start_link(MyServer) - """ - def start_link(opts \\ []) do - server = Keyword.fetch!(opts, :server) - registry = Keyword.get(opts, :registry, Anubis.Server.Registry) - name = registry.supervisor(@kind, server) - - case DynamicSupervisor.start_link(__MODULE__, {server, registry}, name: name) do - {:ok, _pid} = success -> - # Restore sessions from store if configured - restore_sessions(server, registry) - success - - error -> - error - end - end - - @doc """ - Creates a new session for a client connection. - - ## Parameters - * `registry` - The registry module to use to retrieve processes names - * `server` - The server module atom - * `session_id` - Unique identifier for the session (typically from transport) - - ## Returns - * `{:ok, pid}` - Session created successfully - * `{:error, {:already_started, pid}}` - Session already exists - * `{:error, reason}` - Failed to create session - - ## Examples - - # Create a new session for a client - {:ok, session_pid} = Session.Supervisor.create_session(MyRegistry, MyServer, "session-123") - - # Attempting to create duplicate session - {:error, {:already_started, ^session_pid}} = - Session.Supervisor.create_session(MyRegistry, MyServer, "session-123") - """ - def create_session(registry \\ Anubis.Server.Registry, server, session_id) do - name = registry.supervisor(@kind, server) - session_name = registry.server_session(server, session_id) - - DynamicSupervisor.start_child( - name, - {Session, session_id: session_id, name: session_name, server_module: server} - ) - end - - @doc """ - Terminates a session and cleans up its resources. - - ## Parameters - * `registry` - The registry module to use to retrieve processes names - * `server` - The server module atom - * `session_id` - The session identifier to terminate - - ## Returns - * `:ok` - Session terminated successfully - * `{:error, :not_found}` - Session does not exist - - ## Examples - - # Close an existing session - :ok = Session.Supervisor.close_session(MyRegistry, MyServer, "session-123") - - # Attempting to close non-existent session - {:error, :not_found} = Session.Supervisor.close_session(MyRegistry, MyServer, "unknown") - """ - def close_session(registry \\ Anubis.Server.Registry, server, session_id) when is_binary(session_id) do - name = registry.supervisor(@kind, server) - - if pid = registry.whereis_server_session(server, session_id) do - DynamicSupervisor.terminate_child(name, pid) - else - {:error, :not_found} - end - end - - @impl DynamicSupervisor - def init({_server, _registry}) do - DynamicSupervisor.init(strategy: :one_for_one) - end - - # Private functions - - defp restore_sessions(server, registry) do - case Anubis.get_session_store_adapter() do - nil -> - Logging.log(:debug, "No session store configured, skipping session restoration", []) - - store -> - Logging.log(:debug, "Checking for sessions to restore from store", server: server) - - case store.list_active(server: server) do - {:ok, session_ids} -> - Enum.each(session_ids, &create_session(registry, server, &1)) - - {:error, reason} -> - Logging.log(:warning, "Failed to list active sessions from store", server: server, reason: reason) - end - end - end -end diff --git a/lib/anubis/server/supervisor.ex b/lib/anubis/server/supervisor.ex index e2abc35c..164a625b 100644 --- a/lib/anubis/server/supervisor.ex +++ b/lib/anubis/server/supervisor.ex @@ -4,7 +4,7 @@ defmodule Anubis.Server.Supervisor do use Supervisor, restart: :permanent use Anubis.Logging - alias Anubis.Server.Base + alias Anubis.Server.Registry alias Anubis.Server.Session alias Anubis.Server.Transport.SSE alias Anubis.Server.Transport.STDIO @@ -18,9 +18,9 @@ defmodule Anubis.Server.Supervisor do @type start_option :: {:transport, transport} | {:name, Supervisor.name()} + | {:registry, {module(), keyword()}} | {:session_idle_timeout, pos_integer() | nil} | {:request_timeout, pos_integer() | nil} - | {:server_name, GenServer.name() | nil} @doc """ Starts the server supervisor. @@ -30,65 +30,70 @@ defmodule Anubis.Server.Supervisor do * `server` - The module implementing `Anubis.Server` * `opts` - Options including: * `:transport` - Transport configuration (required) - * `:name` - Supervisor name (optional, defaults to registered name) - * `:registry` - The custom registry to use to manage processes names (defaults to `Anubis.Server.Registry`) + * `:name` - Supervisor name (optional, defaults to atom name) + * `:registry` - `{module, opts}` for custom registry (auto-selected by default) * `:session_idle_timeout` - Time in milliseconds before idle sessions expire (default: 30 minutes) - * `:request_timeout` - Time limit in miliseconds for server requests timeout (defaults to 30s) - * `:server_name` - Custom server name, non derived from the `server_module` - - ## Examples - - # Start with STDIO transport - Anubis.Server.Supervisor.start_link(MyServer, [], transport: :stdio) - - # Start with StreamableHTTP transport - Anubis.Server.Supervisor.start_link(MyServer, [], - transport: {:streamable_http, port: 8080} - ) - - # With custom session timeout (15 minutes) - Anubis.Server.Supervisor.start_link(MyServer, [], - transport: {:streamable_http, port: 8080}, - session_idle_timeout: :timer.minutes(15) - ) + * `:request_timeout` - Time limit in milliseconds for server requests (defaults to 30s) """ @spec start_link(server :: module, list(start_option)) :: Supervisor.on_start() def start_link(server, opts) when is_atom(server) and is_list(opts) do - registry = Keyword.get(opts, :registry, Anubis.Server.Registry) - name = Keyword.get(opts, :name, registry.supervisor(server)) - opts = Keyword.merge(opts, module: server, registry: registry) + name = Keyword.get(opts, :name, Registry.supervisor_name(server)) + opts = Keyword.put(opts, :module, server) Supervisor.start_link(__MODULE__, opts, name: name) end + @doc """ + Starts a new session under the DynamicSupervisor. + """ + @spec start_session(module(), keyword()) :: DynamicSupervisor.on_start_child() + def start_session(server, opts) do + sup_name = Registry.session_supervisor_name(server) + DynamicSupervisor.start_child(sup_name, {Session, opts}) + end + + @doc """ + Terminates a session. + """ + @spec stop_session(module(), module(), String.t()) :: :ok | {:error, :not_found} + def stop_session(server, registry_mod, session_id) do + registry_name = Registry.registry_name(server) + + case registry_mod.lookup_session(registry_name, session_id) do + {:ok, pid} -> + sup_name = Registry.session_supervisor_name(server) + DynamicSupervisor.terminate_child(sup_name, pid) + + {:error, :not_found} -> + {:error, :not_found} + end + end + @impl true def init(opts) do server = Keyword.fetch!(opts, :module) transport = normalize_transport(Keyword.fetch!(opts, :transport)) - registry = Keyword.fetch!(opts, :registry) if should_start?(transport) do - {layer, transport_opts} = parse_transport_child(transport, server, registry) - - server_name = registry.server(opts[:server_name] || server) - server_transport = [layer: layer, name: transport_opts[:name]] - - server_opts = [ - module: server, - name: server_name, - transport: server_transport, - registry: registry - ] - - server_opts = - if timeout = Keyword.get(opts, :session_idle_timeout) do - Keyword.put(server_opts, :session_idle_timeout, timeout) - else - server_opts - end - + session_idle_timeout = Keyword.get(opts, :session_idle_timeout) request_timeout = Keyword.get(opts, :request_timeout, to_timeout(second: 30)) + task_supervisor = Registry.task_supervisor_name(server) + + {registry_mod, registry_opts} = resolve_registry(opts, transport, server) + + {layer, transport_opts} = parse_transport_child(transport, server) + + transport_name = transport_opts[:name] - task_supervisor = registry.task_supervisor(server) + session_config = %{ + server_module: server, + registry_mod: registry_mod, + transport: [layer: layer, name: transport_name], + session_idle_timeout: session_idle_timeout, + timeout: request_timeout, + task_supervisor: task_supervisor + } + + :persistent_term.put({__MODULE__, server, :session_config}, session_config) transport_opts = Keyword.merge(transport_opts, @@ -96,12 +101,14 @@ defmodule Anubis.Server.Supervisor do task_supervisor: task_supervisor ) - children = [ - {Task.Supervisor, name: task_supervisor}, - {Session.Supervisor, server: server, registry: registry}, - {Base, server_opts}, - {layer, transport_opts} - ] + children = + case transport do + :stdio -> + build_stdio_children(server, layer, transport_opts, task_supervisor, session_config) + + _ -> + build_http_children(server, registry_mod, registry_opts, layer, transport_opts, task_supervisor) + end Supervisor.init(children, strategy: :one_for_all) else @@ -109,32 +116,99 @@ defmodule Anubis.Server.Supervisor do end end + @doc false + def get_session_config(server) do + :persistent_term.get({__MODULE__, server, :session_config}) + end + + # Auto-select registry: STDIO -> None, HTTP -> Local + defp resolve_registry(opts, transport, server) do + case Keyword.get(opts, :registry) do + {mod, registry_opts} -> + {mod, registry_opts} + + nil -> + case transport do + :stdio -> + {Registry.None, []} + + _ -> + name = Registry.registry_name(server) + {Registry.Local, [name: name]} + end + end + end + + # For STDIO: single session, no DynamicSupervisor, no registry + defp build_stdio_children(server, layer, transport_opts, task_supervisor, session_config) do + session_name = Registry.stdio_session_name(server) + + session_opts = [ + session_id: "stdio", + server_module: server, + name: session_name, + transport: session_config.transport, + session_idle_timeout: session_config.session_idle_timeout || to_timeout(minute: 30), + timeout: session_config.timeout, + task_supervisor: task_supervisor + ] + + [ + {Task.Supervisor, name: task_supervisor}, + {Session, session_opts}, + {layer, transport_opts} + ] + end + + # For HTTP transports: DynamicSupervisor for sessions + registry + defp build_http_children(server, registry_mod, registry_opts, layer, transport_opts, task_supervisor) do + session_sup_name = Registry.session_supervisor_name(server) + + registry_child = + case registry_mod.child_spec(registry_opts) do + :ignore -> nil + spec -> spec + end + + children = [ + {Task.Supervisor, name: task_supervisor}, + {DynamicSupervisor, name: session_sup_name, strategy: :one_for_one}, + {layer, transport_opts} + ] + + if registry_child do + [registry_child | children] + else + children + end + end + defp normalize_transport(t) when t in [:stdio, StubTransport], do: t defp normalize_transport(t) when t in ~w(sse streamable_http)a, do: {t, []} defp normalize_transport({t, opts}) when t in ~w(sse streamable_http)a, do: {t, opts} if Mix.env() == :test do - defp parse_transport_child(StubTransport = kind, server, registry) do - name = registry.transport(server, kind) - opts = [name: name, server: server, registry: registry] + defp parse_transport_child(StubTransport = kind, server) do + name = Registry.transport_name(server, kind) + opts = [name: name, server: server] {kind, opts} end end - defp parse_transport_child(:stdio, server, registry) do - name = registry.transport(server, :stdio) - opts = [name: name, server: server, registry: registry] + defp parse_transport_child(:stdio, server) do + name = Registry.transport_name(server, :stdio) + opts = [name: name, server: server] {STDIO, opts} end - defp parse_transport_child({:streamable_http, opts}, server, registry) do - name = registry.transport(server, :streamable_http) - opts = Keyword.merge(opts, name: name, server: server, registry: registry) + defp parse_transport_child({:streamable_http, opts}, server) do + name = Registry.transport_name(server, :streamable_http) + opts = Keyword.merge(opts, name: name, server: server) {StreamableHTTP, opts} end - defp parse_transport_child({:sse, opts}, server, registry) do + defp parse_transport_child({:sse, opts}, server) do Logging.log( :warning, "The :sse transport option is deprecated as of MCP specification 2025-03-26. " <> @@ -143,8 +217,8 @@ defmodule Anubis.Server.Supervisor do [] ) - name = registry.transport(server, :sse) - opts = Keyword.merge(opts, name: name, server: server, registry: registry) + name = Registry.transport_name(server, :sse) + opts = Keyword.merge(opts, name: name, server: server) {SSE, opts} end diff --git a/lib/anubis/server/transport/sse.ex b/lib/anubis/server/transport/sse.ex index 73dcc99f..e460de4d 100644 --- a/lib/anubis/server/transport/sse.ex +++ b/lib/anubis/server/transport/sse.ex @@ -335,20 +335,15 @@ defmodule Anubis.Server.Transport.SSE do @impl GenServer def handle_call({:handle_message, session_id, message, context}, _from, state) when is_map(message) do - server = state.registry.whereis_server(state.server) - timeout = state.request_timeout - - if Message.is_notification(message) do - GenServer.cast(server, {:notification, message, session_id, context}) - {:reply, {:ok, nil}, state} - else - case forward_request_to_server(server, message, session_id, context, timeout) do - {:ok, response} -> - maybe_send_through_sse(response, session_id, state) - - {:error, reason} -> - {:reply, {:error, reason}, state} - end + case dispatch_session_message(session_id, message, context, state) do + {:ok, response} -> + maybe_send_through_sse(response, session_id, state) + + {:ok_cast} -> + {:reply, {:ok, nil}, state} + + {:error, reason} -> + {:reply, {:error, reason}, state} end end @@ -398,6 +393,22 @@ defmodule Anubis.Server.Transport.SSE do {:reply, endpoint_url, state} end + defp dispatch_session_message(session_id, message, context, state) do + session_pid = state.registry.session_name(state.server, session_id) + + cond do + is_nil(session_pid) -> + {:error, :no_session} + + Message.is_notification(message) -> + GenServer.cast(session_pid, {:mcp_notification, message, context}) + {:ok_cast} + + true -> + forward_request_to_session(session_pid, message, context, state.request_timeout) + end + end + defp maybe_send_through_sse(response, session_id, state) do case Map.get(state.sse_handlers, session_id) do {pid, _ref} -> @@ -409,23 +420,27 @@ defmodule Anubis.Server.Transport.SSE do end end - defp forward_request_to_server(server, message, session_id, context, timeout) do - msg = {:request, message, session_id, context} + defp forward_request_to_session(session_pid, message, context, timeout) do + msg = {:mcp_request, message, context} - case GenServer.call(server, msg, timeout) do + case GenServer.call(session_pid, msg, timeout) do {:ok, response} -> {:ok, response} {:error, reason} -> Logging.transport_event( "server_error", - %{reason: reason, session_id: session_id}, + %{reason: reason}, level: :error ) {:error, reason} end catch + :exit, {:noproc, _} = reason -> + Logging.transport_event("session_not_found", %{reason: reason}, level: :warning) + {:error, :no_session} + :exit, reason -> Logging.transport_event("server_call_failed", %{reason: reason}, level: :error) {:error, :server_unavailable} diff --git a/lib/anubis/server/transport/sse/plug.ex b/lib/anubis/server/transport/sse/plug.ex index e7cbc7f5..854ecdaf 100644 --- a/lib/anubis/server/transport/sse/plug.ex +++ b/lib/anubis/server/transport/sse/plug.ex @@ -91,6 +91,7 @@ if Code.ensure_loaded?(Plug) do alias Anubis.MCP.Error alias Anubis.MCP.ID alias Anubis.MCP.Message + alias Anubis.Server.Registry alias Anubis.Server.Transport.SSE alias Anubis.SSE.Streaming alias Plug.Conn.Unfetched @@ -112,8 +113,7 @@ if Code.ensure_loaded?(Plug) do raise ArgumentError, "SSE.Plug requires :mode to be either :sse or :post" end - registry = Keyword.get(opts, :registry, Anubis.Server.Registry) - transport = registry.transport(server, :sse) + transport = Registry.transport_name(server, :sse) timeout = Keyword.get(opts, :timeout, @default_timeout) %{ diff --git a/lib/anubis/server/transport/stdio.ex b/lib/anubis/server/transport/stdio.ex index 80320eef..f0d6d848 100644 --- a/lib/anubis/server/transport/stdio.ex +++ b/lib/anubis/server/transport/stdio.ex @@ -3,7 +3,8 @@ defmodule Anubis.Server.Transport.STDIO do STDIO transport implementation for MCP servers. This module handles communication with MCP clients via standard input/output streams, - processing incoming JSON-RPC messages and forwarding responses. + processing incoming JSON-RPC messages and forwarding responses directly to the + Session process. """ @behaviour Anubis.Transport @@ -15,6 +16,7 @@ defmodule Anubis.Server.Transport.STDIO do import Peri alias Anubis.MCP.Message + alias Anubis.Server.Registry alias Anubis.Telemetry alias Anubis.Transport.Behaviour, as: Transport @@ -95,7 +97,7 @@ defmodule Anubis.Server.Transport.STDIO do @typedoc """ STDIO transport options - - `:server` - The server process (required) + - `:server` - The server module (required) - `:name` - Optional name for registering the GenServer """ @type option :: @@ -106,23 +108,9 @@ defmodule Anubis.Server.Transport.STDIO do defschema(:parse_options, [ {:server, {:required, {:oneof, [{:custom, &Anubis.genserver_name/1}, :pid, {:tuple, [:atom, :any]}]}}}, {:name, {:custom, &Anubis.genserver_name/1}}, - {:registry, {:atom, {:default, Anubis.Server.Registry}}}, {:request_timeout, {:integer, {:default, to_timeout(second: 30)}}} ]) - @doc """ - Starts a new STDIO transport process. - - ## Parameters - * `opts` - Options - * `:server` - (required) The server to forward messages to - * `:name` - Optional name for the GenServer process - - ## Examples - - iex> Anubis.Server.Transport.STDIO.start_link(server: my_server) - {:ok, pid} - """ @impl Transport @spec start_link(Enumerable.t(option())) :: GenServer.on_start() def start_link(opts) do @@ -136,28 +124,11 @@ defmodule Anubis.Server.Transport.STDIO do end end - @doc """ - Sends a message to the client via stdout. - - ## Parameters - * `transport` - The transport process - * `message` - The message to send - - ## Returns - * `:ok` if message was sent successfully - * `{:error, reason}` otherwise - """ @impl Transport def send_message(transport, message, opts) when is_binary(message) do GenServer.call(transport, {:send, message}, opts[:timeout]) end - @doc """ - Shuts down the transport connection. - - ## Parameters - * `transport` - The transport process - """ @impl Transport @spec shutdown(GenServer.server()) :: :ok def shutdown(transport) do @@ -175,7 +146,6 @@ defmodule Anubis.Server.Transport.STDIO do state = %{ server: opts.server, reading_task: nil, - registry: opts.registry, request_timeout: opts.request_timeout } @@ -219,7 +189,7 @@ defmodule Anubis.Server.Transport.STDIO do end @impl GenServer - def handle_cast({:send, message}, state) do + def handle_call({:send, message}, _from, state) do Logging.transport_event( "outgoing", %{transport: :stdio, message_size: byte_size(message)}, @@ -233,7 +203,7 @@ defmodule Anubis.Server.Transport.STDIO do ) IO.write(message) - {:noreply, state} + {:reply, :ok, state} end @impl GenServer @@ -317,9 +287,8 @@ defmodule Anubis.Server.Transport.STDIO do end end - defp process_message(message, %{server: server_name, registry: registry} = state) do - server = registry.whereis_server(server_name) - timeout = state.request_timeout + defp process_message(message, %{server: server_module} = state) do + session_pid = Registry.stdio_session_name(server_module) context = %{ type: :stdio, @@ -327,21 +296,40 @@ defmodule Anubis.Server.Transport.STDIO do pid: System.pid() } + case get_session_pid(session_pid) do + {:ok, pid} -> + dispatch_to_session(message, pid, context, state) + + :error -> + Logging.transport_event("no_session", %{server: server_module}, level: :error) + end + end + + defp get_session_pid(session_name) do + if Process.whereis(session_name), do: {:ok, session_name}, else: :error + end + + defp dispatch_to_session(message, session_pid, context, state) do if Message.is_notification(message) do - GenServer.cast(server, {:notification, message, "stdio", context}) + GenServer.cast(session_pid, {:mcp_notification, message, context}) else - case GenServer.call(server, {:request, message, "stdio", context}, timeout) do - {:ok, response} when is_binary(response) -> - # send_message(self(), response) - # NOTE: will be fixed soon, we need to rewrite stdio for server - :ok - - {:error, reason} -> - Logging.transport_event("server_error", %{reason: reason}, level: :error) - end + forward_request_to_session(session_pid, message, context, state.request_timeout) + end + end + + defp forward_request_to_session(session_pid, message, context, timeout) do + case GenServer.call(session_pid, {:mcp_request, message, context}, timeout) do + {:ok, response} when is_binary(response) -> + IO.write(response <> "\n") + + {:ok, nil} -> + :ok + + {:error, reason} -> + Logging.transport_event("session_error", %{reason: reason}, level: :error) end catch :exit, reason -> - Logging.transport_event("server_call_failed", %{reason: reason}, level: :error) + Logging.transport_event("session_call_failed", %{reason: reason}, level: :error) end end diff --git a/lib/anubis/server/transport/streamable_http.ex b/lib/anubis/server/transport/streamable_http.ex index dd674787..916c4c9b 100644 --- a/lib/anubis/server/transport/streamable_http.ex +++ b/lib/anubis/server/transport/streamable_http.ex @@ -2,18 +2,16 @@ defmodule Anubis.Server.Transport.StreamableHTTP do @moduledoc """ StreamableHTTP transport implementation for MCP servers. - This module provides an HTTP-based transport layer that supports multiple - concurrent client sessions through Server-Sent Events (SSE). It enables - web-based MCP clients to communicate with the server using standard HTTP - protocols. + This module manages SSE (Server-Sent Events) connections for server-to-client + communication. In the refactored architecture, request handling is done directly + by Session processes - this module only manages SSE handlers and notifications. ## Features - - Multiple concurrent client sessions - - Server-Sent Events for real-time server-to-client communication - - HTTP POST endpoint for client-to-server messages - - Automatic session cleanup on disconnect - - Integration with Phoenix/Plug applications + - SSE handler registration for server-to-client push + - Automatic handler cleanup on disconnect + - Keepalive messages to maintain connections + - Notification broadcasting to connected clients ## Usage @@ -29,19 +27,6 @@ defmodule Anubis.Server.Transport.StreamableHTTP do # In your router forward "/mcp", Anubis.Server.Transport.StreamableHTTP.Plug, server: MyApp.MCPServer - - ## Message Flow - - 1. Client connects to `/sse` endpoint, receives a session ID - 2. Client sends messages via POST to `/messages` with session ID header - 3. Server responses are pushed through the SSE connection - 4. Connection closes on client disconnect or server shutdown - - ## Configuration - - - `:port` - HTTP server port (default: 4000) - - `:server` - The MCP server process to connect to - - `:name` - Process registration name """ @behaviour Anubis.Transport @@ -52,15 +37,9 @@ defmodule Anubis.Server.Transport.StreamableHTTP do import Peri - alias Anubis.MCP.Error - alias Anubis.MCP.ID - alias Anubis.MCP.Message - alias Anubis.Server.Transport.StreamableHTTP.RequestParams alias Anubis.Telemetry alias Anubis.Transport.Behaviour, as: Transport - require Message - @type t :: GenServer.server() @type http_state :: %{ @@ -157,19 +136,10 @@ defmodule Anubis.Server.Transport.StreamableHTTP do end) end - @type request_t :: %RequestParams{ - transport: GenServer.server(), - session_id: String.t() | nil, - session_header: String.t(), - timeout: pos_integer(), - context: map() | nil, - message: map() | binary() | nil - } - @typedoc """ StreamableHTTP transport options - - `:server` - The server process (required) + - `:server` - The server module (required) - `:name` - Name for registering the GenServer (required) """ @type option :: @@ -186,9 +156,6 @@ defmodule Anubis.Server.Transport.StreamableHTTP do {:keepalive_interval, {:integer, {:default, 5_000}}} ]) - @doc """ - Starts the StreamableHTTP transport. - """ @impl Transport @spec start_link(Enumerable.t(option())) :: GenServer.on_start() def start_link(opts) do @@ -198,33 +165,11 @@ defmodule Anubis.Server.Transport.StreamableHTTP do GenServer.start_link(__MODULE__, Map.new(opts), name: name) end - @doc """ - Sends a message to the client via the active SSE connection. - - This function is used for server-initiated notifications. - It will broadcast to all active SSE connections. - - ## Parameters - * `transport` - The transport process - * `message` - The message to send - - ## Returns - * `:ok` if message was sent successfully - * `{:error, reason}` otherwise - """ @impl Transport def send_message(transport, message, opts) when is_binary(message) do GenServer.call(transport, {:send_message, message}, opts[:timeout]) end - @doc """ - Shuts down the transport connection. - - This terminates all active sessions managed by this transport. - - ## Parameters - * `transport` - The transport process - """ @impl Transport @spec shutdown(GenServer.server()) :: :ok def shutdown(transport) do @@ -235,10 +180,9 @@ defmodule Anubis.Server.Transport.StreamableHTTP do def supported_protocol_versions, do: ["2025-03-26", "2025-06-18"] @doc """ - Registers an SSE handler process for a session. + Registers the calling process as the SSE handler for a session. Called by the Plug when establishing an SSE connection. - The calling process becomes the SSE handler for the session. """ @spec register_sse_handler(GenServer.server(), String.t()) :: :ok | {:error, term()} def register_sse_handler(transport, session_id) do @@ -246,9 +190,7 @@ defmodule Anubis.Server.Transport.StreamableHTTP do end @doc """ - Unregisters an SSE handler process for a session. - - Called when the SSE connection is closed. + Unregisters the SSE handler for a session. Called when the SSE connection closes. """ @spec unregister_sse_handler(GenServer.server(), String.t()) :: :ok def unregister_sse_handler(transport, session_id) do @@ -256,34 +198,7 @@ defmodule Anubis.Server.Transport.StreamableHTTP do end @doc """ - Handles an incoming message from a client with request context. - - Called by the Plug when a message is received via HTTP POST. - """ - @spec handle_message(request_t) :: {:ok, binary() | nil} | {:error, term()} - def handle_message(%RequestParams{transport: transport} = params) do - timeout = params.timeout + 1_000 - GenServer.call(transport, {:handle_message, params}, timeout) - end - - @doc """ - Handles an incoming message with context and returns {:sse, response} if SSE handler exists. - - This allows the Plug to know whether to stream the response via SSE - or return it as a regular HTTP response. - """ - @spec handle_message_for_sse(request_t) :: - {:ok, binary()} | {:sse, binary()} | {:error, term()} - def handle_message_for_sse(%RequestParams{transport: transport} = params) do - timeout = params.timeout + 1_000 - GenServer.call(transport, {:handle_message_for_sse, params}, timeout) - end - - @doc """ - Gets the SSE handler process for a session. - - Returns the pid of the process handling SSE for this session, - or nil if no SSE connection exists. + Returns the SSE handler pid for a session, or `nil` if none is connected. """ @spec get_sse_handler(GenServer.server(), String.t()) :: pid() | nil def get_sse_handler(transport, session_id) do @@ -291,9 +206,7 @@ defmodule Anubis.Server.Transport.StreamableHTTP do end @doc """ - Routes a message to a specific session's SSE handler. - - Used for targeted server notifications to specific clients. + Routes a message to a specific session's SSE handler for server-to-client push. """ @spec route_to_session(GenServer.server(), String.t(), binary()) :: :ok | {:error, term()} @@ -311,10 +224,7 @@ defmodule Anubis.Server.Transport.StreamableHTTP do server: server, registry: opts.registry, task_supervisor: opts.task_supervisor, - # Map of session_id => {pid, monitor_ref} sse_handlers: %{}, - active_tasks: %{}, - # keepalive keepalive_interval: opts.keepalive_interval, keepalive_enabled: opts.keepalive } @@ -353,71 +263,6 @@ defmodule Anubis.Server.Transport.StreamableHTTP do {:reply, :ok, %{state | sse_handlers: sse_handlers}} end - @impl GenServer - def handle_call({:handle_message, %{message: message} = params}, from, state) when is_map(message) do - %{session_id: session_id, context: context, timeout: timeout} = params - server = state.registry.whereis_server(state.server) - - cond do - Message.is_notification(params.message) -> - GenServer.cast(server, {:notification, message, session_id, context}) - {:reply, {:ok, nil}, state} - - Message.is_response(message) or Message.is_error(message) -> - GenServer.cast(server, {:response, message, session_id, context}) - {:reply, {:ok, nil}, state} - - true -> - task = - Task.Supervisor.async_nolink(state.task_supervisor, fn -> - forward_request_to_server(server, params) - end) - - task_timeout_ref = Process.send_after(self(), {:task_timeout, task.ref}, timeout) - - task_info = %{ - type: :handle_message, - session_id: session_id, - from: from, - task_timeout: task_timeout_ref, - task: task - } - - {:noreply, put_in(state.active_tasks[task.ref], task_info)} - end - end - - @impl GenServer - def handle_call({:handle_message_for_sse, %{message: message} = params}, from, state) when is_map(message) do - %{session_id: session_id, context: context, timeout: timeout} = params - server = state.registry.whereis_server(state.server) - - if Message.is_notification(message) do - GenServer.cast(server, {:notification, message, session_id, context}) - {:reply, {:ok, nil}, state} - else - sse_handler? = Map.has_key?(state.sse_handlers, session_id) - - task = - Task.Supervisor.async_nolink(state.task_supervisor, fn -> - forward_request_to_server(server, params, sse_handler?) - end) - - task_timeout_ref = Process.send_after(self(), {:task_timeout, task.ref}, timeout) - - task_info = %{ - type: :handle_message_for_sse, - session_id: session_id, - from: from, - has_sse_handler: sse_handler?, - task_timeout: task_timeout_ref, - task: task - } - - {:noreply, put_in(state.active_tasks[task.ref], task_info)} - end - end - @impl GenServer def handle_call({:get_sse_handler, session_id}, _from, state) do case Map.get(state.sse_handlers, session_id) do @@ -452,31 +297,6 @@ defmodule Anubis.Server.Transport.StreamableHTTP do {:reply, :ok, state} end - defp forward_request_to_server(server, params, has_sse_handler \\ false) do - msg = {:request, params.message, params.session_id, params.context} - - case GenServer.call(server, msg, params.timeout) do - {:ok, response} when has_sse_handler -> - {:sse, response} - - {:ok, response} -> - {:ok, response} - - {:error, reason} -> - Logging.transport_event( - "server_error", - %{reason: reason, session_id: params.session_id}, - level: :error - ) - - {:error, reason} - end - catch - :exit, reason -> - Logging.transport_event("server_call_failed", %{reason: reason}, level: :error) - {:error, :server_unavailable} - end - @impl GenServer def handle_cast({:unregister_sse_handler, session_id}, state) do sse_handlers = @@ -509,69 +329,7 @@ defmodule Anubis.Server.Transport.StreamableHTTP do {:stop, :normal, state} end - # Handle successful task completion @impl GenServer - def handle_info({:task_timeout, ref}, %{active_tasks: active_tasks} = state) when is_map_key(active_tasks, ref) do - {task_info, active_tasks} = Map.pop(active_tasks, ref) - - timeout_error = - Error.protocol(:internal_error, %{ - message: "Request timeout - tool execution exceeded limit", - session_id: task_info.session_id - }) - - {:ok, error_json} = Error.to_json_rpc(timeout_error, ID.generate_error_id()) - - GenServer.reply(task_info.from, {:error, error_json}) - if task = task_info.task, do: Task.shutdown(task, :brutal_kill) - - Logging.transport_event( - "task_timeout", - %{ - session_id: task_info.session_id - }, - level: :warning - ) - - {:noreply, %{state | active_tasks: active_tasks}} - end - - def handle_info({:task_timeout, _ref}, state), do: {:noreply, state} - - def handle_info({ref, result}, %{active_tasks: active_tasks} = state) - when is_reference(ref) and is_map_key(active_tasks, ref) do - {task_info, active_tasks} = Map.pop(active_tasks, ref) - - if Map.has_key?(task_info, :timeout_ref) do - Process.cancel_timer(task_info.task_timeout) - end - - GenServer.reply(task_info.from, result) - Process.demonitor(ref, [:flush]) - - {:noreply, %{state | active_tasks: active_tasks}} - end - - def handle_info({_ref, _}, state), do: {:noreply, state} - - def handle_info({:DOWN, ref, :process, _pid, reason}, %{active_tasks: active_tasks} = state) - when is_map_key(active_tasks, ref) do - {task_info, active_tasks} = Map.pop(active_tasks, ref) - error = {:error, {:task_crashed, reason}} - GenServer.reply(task_info.from, error) - - Logging.transport_event( - "task_crashed", - %{ - reason: inspect(reason, pretty: true), - session_id: task_info.session_id - }, - level: :error - ) - - {:noreply, %{state | active_tasks: active_tasks}} - end - def handle_info({:DOWN, ref, :process, pid, reason}, state) do sse_handlers = state.sse_handlers diff --git a/lib/anubis/server/transport/streamable_http/plug.ex b/lib/anubis/server/transport/streamable_http/plug.ex index 5ac3ee96..a72cc901 100644 --- a/lib/anubis/server/transport/streamable_http/plug.ex +++ b/lib/anubis/server/transport/streamable_http/plug.ex @@ -10,11 +10,6 @@ if Code.ensure_loaded?(Plug) do - POST: Handles JSON-RPC messages from client to server - DELETE: Closes a session - ## SSE Streaming Architecture - - This Plug handles SSE streaming by keeping the request process alive - and managing the streaming loop for server-to-client communication. - ## Usage in Phoenix Router pipeline :mcp do @@ -26,32 +21,11 @@ if Code.ensure_loaded?(Plug) do forward "/", to: Anubis.Server.Transport.StreamableHTTP.Plug, server: :your_server_name end - ## Usage in Plug Router - - forward "/mcp", to: Anubis.Server.Transport.StreamableHTTP.Plug, init_opts: [server: :your_server_name] - ## Configuration Options - `:server` - The server process name (required) - `:session_header` - Custom header name for session ID (default: "mcp-session-id") - `:request_timeout` - Request timeout in milliseconds (default: 30000) - - `:registry` - The registry to use. See `Anubis.Server.Registry.Adapter` for more information (default: Elixir's Registry implementation) - - ## Security Features - - - Origin header validation for DNS rebinding protection - - Session-based request validation - - Automatic session cleanup on connection loss - - Rate limiting support (when configured) - - ## HTTP Response Codes - - - 200: Successful request - - 202: Accepted (for notifications and responses) - - 400: Bad request (malformed JSON-RPC) - - 404: Session not found - - 405: Method not allowed - - 500: Internal server error """ @behaviour Plug @@ -63,8 +37,9 @@ if Code.ensure_loaded?(Plug) do alias Anubis.MCP.Error alias Anubis.MCP.ID alias Anubis.MCP.Message + alias Anubis.Server.Registry + alias Anubis.Server.Supervisor, as: ServerSupervisor alias Anubis.Server.Transport.StreamableHTTP - alias Anubis.Server.Transport.StreamableHTTP.RequestParams alias Anubis.SSE.Streaming alias Plug.Conn.Unfetched @@ -78,15 +53,16 @@ if Code.ensure_loaded?(Plug) do @impl Plug def init(opts) do server = Keyword.fetch!(opts, :server) - registry = Keyword.get(opts, :registry, Anubis.Server.Registry) - transport = registry.transport(server, :streamable_http) + session_config = ServerSupervisor.get_session_config(server) + transport_name = Registry.transport_name(server, :streamable_http) session_header = Keyword.get(opts, :session_header, @default_session_header) request_timeout = Keyword.get(opts, :request_timeout, @default_timeout) %{ server: server, - registry: registry, - transport: transport, + registry_mod: session_config.registry_mod, + registry_name: Registry.registry_name(server), + transport: transport_name, session_header: session_header, timeout: request_timeout } @@ -122,9 +98,9 @@ if Code.ensure_loaded?(Plug) do end end - # POST request handler - processes MCP messages + # POST request handler - processes MCP messages directly to Session - defp handle_post(conn, %{transport: transport, session_header: session_header} = opts) do + defp handle_post(conn, %{session_header: session_header} = opts) do with :ok <- validate_accept_header(conn), {:ok, body, conn} <- maybe_read_request_body(conn, opts), {:ok, [message]} <- maybe_parse_messages(body) do @@ -136,17 +112,7 @@ if Code.ensure_loaded?(Plug) do session_id: session_id }) - process_message( - conn, - RequestParams.new( - message: message, - transport: transport, - session_id: session_id, - context: context, - session_header: session_header, - timeout: opts.timeout - ) - ) + process_message(conn, message, session_id, context, opts) else {:error, :invalid_accept_header} -> send_error( @@ -173,136 +139,166 @@ if Code.ensure_loaded?(Plug) do end end - defp process_message(conn, %{message: message} = params) when is_map(message) do - if Message.is_request(message) do - handle_request_with_possible_sse(conn, params) - else - # Notification - params - |> StreamableHTTP.handle_message() - |> format_notification_response(conn) - end - end + defp process_message(conn, message, session_id, context, opts) do + cond do + Message.is_notification(message) -> + handle_notification_message(conn, message, session_id, context, opts) - defp format_notification_response({:ok, _}, conn) do - conn - |> put_resp_content_type("application/json") - |> send_resp(202, "{}") - end + Message.is_response(message) or Message.is_error(message) -> + handle_response_message(conn, message, session_id, context, opts) - defp format_notification_response({:error, %Error{} = error}, conn) do - send_jsonrpc_error(conn, error, nil) + Message.is_request(message) -> + handle_request_message(conn, message, session_id, context, opts) + + true -> + send_jsonrpc_error( + conn, + Error.protocol(:invalid_request, %{message: "Invalid message type"}), + nil + ) + end end - defp format_notification_response({:error, reason}, conn) do - Logging.transport_event("notification_handling_failed", %{reason: reason}, level: :error) + defp handle_notification_message(conn, message, session_id, context, opts) do + case find_session(opts, session_id) do + {:ok, session_pid} -> + GenServer.cast(session_pid, {:mcp_notification, message, context}) - send_jsonrpc_error( - conn, - Error.protocol(:internal_error, %{reason: reason}), - nil - ) + conn + |> put_resp_content_type("application/json") + |> send_resp(202, "{}") + + {:error, :not_found} -> + send_error(conn, 400, "No active session") + end end - defp handle_delete(conn, %{transport: transport, session_header: session_header} = opts) do - case get_req_header(conn, session_header) do - [session_id] when is_binary(session_id) and session_id != "" -> - StreamableHTTP.unregister_sse_handler(transport, session_id) - delete_session_from_store(session_id) - stop_session_process(opts, session_id) + defp handle_response_message(conn, message, session_id, context, opts) do + case find_session(opts, session_id) do + {:ok, session_pid} -> + GenServer.cast(session_pid, {:mcp_response, message, context}) conn |> put_resp_content_type("application/json") - |> send_resp(200, "{}") + |> send_resp(202, "{}") - _ -> - send_error(conn, 400, "Session ID required") + {:error, :not_found} -> + send_error(conn, 400, "No active session") end end - # Handle requests that might need SSE streaming + defp handle_request_message(conn, message, session_id, context, opts) do + case find_or_create_session(opts, session_id, message) do + {:ok, session_pid} -> + if wants_sse?(conn) do + handle_sse_request(conn, session_pid, message, session_id, context, opts) + else + handle_json_request(conn, session_pid, message, session_id, context, opts) + end - defp handle_request_with_possible_sse(conn, params) do - if wants_sse?(conn) do - handle_sse_request(conn, params) - else - handle_json_request(conn, params) + {:error, :no_session} -> + send_error(conn, 400, "No active session") + + {:error, reason} -> + send_jsonrpc_error( + conn, + Error.protocol(:internal_error, %{reason: reason}), + extract_request_id(message) + ) end end - defp handle_sse_request(conn, params) do - case StreamableHTTP.handle_message_for_sse(params) do - {:sse, response} -> - route_sse_response(conn, response, params) - - {:ok, response} -> + defp handle_json_request(conn, session_pid, message, session_id, context, %{session_header: session_header} = opts) do + case GenServer.call(session_pid, {:mcp_request, message, context}, opts.timeout) do + {:ok, response} when is_binary(response) -> conn |> put_resp_content_type("application/json") - |> maybe_add_session_header(params.session_header, params.session_id) + |> maybe_add_session_header(session_header, session_id) |> send_resp(200, response) + {:ok, nil} -> + conn + |> put_resp_content_type("application/json") + |> maybe_add_session_header(session_header, session_id) + |> send_resp(200, "{}") + {:error, error} -> - handle_request_error(conn, error, params.message) + handle_request_error(conn, error, message) end + catch + :exit, reason -> + Logging.transport_event("session_call_failed", %{reason: reason}, level: :error) + + send_jsonrpc_error( + conn, + Error.protocol(:internal_error, %{message: "Server unavailable"}), + extract_request_id(message) + ) end - defp handle_json_request(conn, params) do - case StreamableHTTP.handle_message(params) do - {:ok, response} -> + defp handle_sse_request(conn, session_pid, message, session_id, context, opts) do + %{session_header: session_header} = opts + + case GenServer.call(session_pid, {:mcp_request, message, context}, opts.timeout) do + {:ok, response} when is_binary(response) -> + route_sse_response(conn, response, session_id, opts) + + {:ok, nil} -> conn |> put_resp_content_type("application/json") - |> maybe_add_session_header(params.session_header, params.session_id) - |> send_resp(200, response) + |> maybe_add_session_header(session_header, session_id) + |> send_resp(200, "{}") {:error, error} -> - handle_request_error(conn, error, params.message) + handle_request_error(conn, error, message) end + catch + :exit, reason -> + Logging.transport_event("session_call_failed", %{reason: reason}, level: :error) + + send_jsonrpc_error( + conn, + Error.protocol(:internal_error, %{message: "Server unavailable"}), + extract_request_id(message) + ) end - defp route_sse_response(conn, response, params) do - %{transport: transport, session_id: session_id} = params + defp route_sse_response(conn, response, session_id, %{transport: transport} = opts) do handler_pid = StreamableHTTP.get_sse_handler(transport, session_id) - # Check Process.alive? to handle race condition where handler died but - # transport hasn't processed the :DOWN message yet. Without this check, - # send/2 silently drops the message to a dead process. - if handler_pid && Process.alive?(handler_pid) do - send(handler_pid, {:sse_message, response}) + cond do + handler_pid && Process.alive?(handler_pid) -> + send(handler_pid, {:sse_message, response}) - conn - |> put_resp_content_type("application/json") - |> send_resp(202, "{}") - else - # Handler is missing or stale - clean up stale entry if present - if handler_pid do + conn + |> put_resp_content_type("application/json") + |> send_resp(202, "{}") + + handler_pid -> StreamableHTTP.unregister_sse_handler(transport, session_id) - end + establish_sse_for_request(conn, response, session_id, opts) - establish_sse_for_request(conn, params) + true -> + establish_sse_for_request(conn, response, session_id, opts) end end - defp handle_request_error(conn, %Error{} = error, body) do - send_jsonrpc_error(conn, error, extract_request_id(body)) - end - - defp handle_request_error(conn, reason, body) do - Logging.transport_event("request_error", %{reason: reason}, level: :error) - - send_jsonrpc_error( - conn, - Error.protocol(:internal_error, %{reason: reason}), - extract_request_id(body) - ) - end - - defp establish_sse_for_request(conn, params) do - %{transport: transport, session_id: session_id} = params + defp establish_sse_for_request(conn, response, session_id, opts) do + %{transport: transport, session_header: session_header} = opts case StreamableHTTP.register_sse_handler(transport, session_id) do :ok -> - start_background_request(params) - start_sse_streaming(conn, params) + self_pid = self() + Task.start(fn -> send(self_pid, {:sse_message, response}) end) + + conn + |> put_resp_header(session_header, session_id) + |> Streaming.prepare_connection() + |> Streaming.start(transport, session_id, + on_close: fn -> + StreamableHTTP.unregister_sse_handler(transport, session_id) + end + ) {:error, reason} -> Logging.transport_event("sse_registration_failed", %{reason: reason}, level: :error) @@ -310,40 +306,71 @@ if Code.ensure_loaded?(Plug) do send_jsonrpc_error( conn, Error.protocol(:internal_error, %{reason: reason}), - extract_request_id(params.message) + nil ) end end - defp start_background_request(params) do - self_pid = self() + defp handle_delete(conn, %{transport: transport, session_header: session_header} = opts) do + case get_req_header(conn, session_header) do + [session_id] when is_binary(session_id) and session_id != "" -> + StreamableHTTP.unregister_sse_handler(transport, session_id) + delete_session_from_store(session_id) + stop_session_process(opts, session_id) - Task.start(fn -> - case StreamableHTTP.handle_message(params) do - {:ok, response} when is_binary(response) -> - send(self_pid, {:sse_message, response}) + conn + |> put_resp_content_type("application/json") + |> send_resp(200, "{}") - {:error, reason} -> - Logging.transport_event( - "sse_background_request_error", - %{reason: reason}, - level: :error - ) - end - end) + _ -> + send_error(conn, 400, "Session ID required") + end end - defp start_sse_streaming(conn, params) do - %{transport: transport, session_id: session_id} = params + # Session management - conn - |> put_resp_header(params.session_header, session_id) - |> Streaming.prepare_connection() - |> Streaming.start(transport, session_id, - on_close: fn -> - StreamableHTTP.unregister_sse_handler(transport, session_id) - end - ) + defp find_session(%{registry_mod: mod, registry_name: name}, session_id) do + mod.lookup_session(name, session_id) + end + + defp find_or_create_session(opts, session_id, message) do + case find_session(opts, session_id) do + {:ok, pid} -> + {:ok, pid} + + {:error, :not_found} when Message.is_initialize(message) -> + start_new_session(opts, session_id) + + {:error, :not_found} -> + {:error, :no_session} + end + end + + defp start_new_session(%{server: server, registry_mod: registry_mod, registry_name: registry_name} = opts, session_id) do + session_config = ServerSupervisor.get_session_config(server) + session_name = Registry.session_name(server, session_id) + + session_opts = [ + session_id: session_id, + server_module: server, + name: session_name, + transport: session_config.transport, + session_idle_timeout: session_config.session_idle_timeout || 1_800_000, + timeout: opts.timeout, + task_supervisor: session_config.task_supervisor + ] + + case ServerSupervisor.start_session(server, session_opts) do + {:ok, pid} -> + registry_mod.register_session(registry_name, session_id, pid) + {:ok, pid} + + {:error, {:already_started, pid}} -> + {:ok, pid} + + {:error, reason} -> + {:error, reason} + end end # Helper functions @@ -361,8 +388,6 @@ if Code.ensure_loaded?(Plug) do |> get_req_header("accept") |> List.first("") - # For POST requests, client must accept application/json at minimum - # text/event-stream is optional and indicates client wants SSE responses if String.contains?(accept_header, "application/json") do :ok else @@ -381,14 +406,11 @@ if Code.ensure_loaded?(Plug) do end defp determine_session_id(conn, session_header, message) when Message.is_initialize(message) do - # For initialize messages, check if client provided a session ID to resume case get_req_header(conn, session_header) do [session_id] when is_binary(session_id) and session_id != "" -> - # Client wants to resume existing session - use their ID session_id _ -> - # No session ID provided - generate new one for fresh session ID.generate_session_id() end end @@ -463,6 +485,20 @@ if Code.ensure_loaded?(Plug) do |> send_resp(400, encoded_error) end + defp handle_request_error(conn, %Error{} = error, body) do + send_jsonrpc_error(conn, error, extract_request_id(body)) + end + + defp handle_request_error(conn, reason, body) do + Logging.transport_event("request_error", %{reason: reason}, level: :error) + + send_jsonrpc_error( + conn, + Error.protocol(:internal_error, %{reason: reason}), + extract_request_id(body) + ) + end + defp extract_request_id(%{"id" => request_id}), do: request_id defp extract_request_id(_), do: nil @@ -487,18 +523,27 @@ if Code.ensure_loaded?(Plug) do end end + defp start_sse_streaming(conn, params) do + %{transport: transport, session_id: session_id, session_header: session_header} = params + + conn + |> put_resp_header(session_header, session_id) + |> Streaming.prepare_connection() + |> Streaming.start(transport, session_id, + on_close: fn -> + StreamableHTTP.unregister_sse_handler(transport, session_id) + end + ) + end + defp delete_session_from_store(session_id) do if store = Anubis.get_session_store_adapter() do store.delete(session_id, []) end end - defp stop_session_process(%{server: server, registry: registry}, session_id) do - session_name = registry.server_session(server, session_id) - - if pid = GenServer.whereis(session_name) do - GenServer.stop(pid, :normal) - end + defp stop_session_process(%{server: server, registry_mod: registry_mod}, session_id) do + ServerSupervisor.stop_session(server, registry_mod, session_id) end end end diff --git a/test/anubis/server/base_test.exs b/test/anubis/server/base_test.exs deleted file mode 100644 index 7e3188ab..00000000 --- a/test/anubis/server/base_test.exs +++ /dev/null @@ -1,364 +0,0 @@ -defmodule Anubis.Server.BaseTest do - use Anubis.MCP.Case, async: false - - alias Anubis.MCP.Message - alias Anubis.Server.Base - alias Anubis.Server.Frame - alias Anubis.Server.Session - - require Message - - @moduletag capture_log: true - - describe "start_link/1" do - test "starts a server with valid options" do - transport = start_supervised!(StubTransport) - - assert {:ok, pid} = - Base.start_link( - module: StubServer, - name: :named_server, - transport: [layer: StubTransport, name: transport] - ) - - assert Process.alive?(pid) - end - - test "starts a named server" do - transport = start_supervised!({StubTransport, []}, id: :named_transport) - - assert {:ok, _pid} = - Base.start_link( - module: StubServer, - name: :named_server, - transport: [layer: StubTransport, name: transport] - ) - - assert pid = Process.whereis(:named_server) - assert Process.alive?(pid) - end - end - - describe "handle_call/3 for messages" do - setup :initialized_server - - @tag skip: true - test "handles errors", %{server: server} do - error = build_error(-32_000, "got wrong", 1) - assert {:ok, _} = GenServer.call(server, {:request, error, "123", %{}}) - end - - test "rejects requests when not initialized", %{server: server} do - request = build_request("tools/list", 123) - - assert {:ok, _} = - GenServer.call(server, {:request, request, "not_initialized", %{}}) - end - - test "accept ping requests when not initialized", %{ - server: server, - session_id: session_id - } do - request = build_request("ping", 123) - assert {:ok, _} = GenServer.call(server, {:request, request, session_id, %{}}) - end - end - - describe "handle_cast/2 for notifications" do - setup :initialized_server - - test "handles notifications", %{server: server, session_id: session_id} do - notification = - build_notification("notifications/cancelled", %{"requestId" => 1}) - - assert :ok = - GenServer.cast(server, {:notification, notification, session_id, %{}}) - end - - test "handles initialize notification", %{server: server, session_id: session_id} do - notification = build_notification("notifications/initialized", %{}) - - assert :ok = - GenServer.cast(server, {:notification, notification, session_id, %{}}) - end - end - - describe "send_notification/3" do - setup :initialized_server - - test "sends notification to transport", ctx do - frame = Frame.put_private(%Frame{}, ctx) - assert :ok = Anubis.Server.send_log_message(frame, :info, "hello") - end - end - - describe "session expiration" do - setup do - start_supervised!(Anubis.Server.Registry) - - start_supervised!({Session.Supervisor, server: StubServer, registry: Anubis.Server.Registry}) - - :ok - end - - test "session expires after idle timeout" do - transport = start_supervised!(StubTransport) - - server = - start_supervised!({Base, - [ - module: StubServer, - name: :expiry_test_server, - transport: [layer: StubTransport, name: transport], - # 100ms for testing - session_idle_timeout: 100 - ]}) - - session_id = "test_session_#{System.unique_integer()}" - - init_msg = - init_request("2025-03-26", %{"name" => "TestClient", "version" => "1.0.0"}) - - assert {:ok, _} = GenServer.call(server, {:request, init_msg, session_id, %{}}) - - init_notification = build_notification("notifications/initialized", %{}) - - assert :ok = - GenServer.cast( - server, - {:notification, init_notification, session_id, %{}} - ) - - session_name = Anubis.Server.Registry.server_session(StubServer, session_id) - assert Session.get(session_name) - - Process.sleep(150) - - # After expiration, the session should no longer be accessible - # The session process has been terminated by the supervisor - assert catch_exit(Session.get(session_name)) - end - - test "session timer resets on activity" do - transport = start_supervised!(StubTransport) - - server = - start_supervised!( - {Base, - [ - module: StubServer, - name: :reset_test_server, - transport: [layer: StubTransport, name: transport], - session_idle_timeout: 200 - ]} - ) - - session_id = "reset_session_#{System.unique_integer()}" - - init_msg = - init_request("2025-03-26", %{"name" => "TestClient", "version" => "1.0.0"}) - - assert {:ok, _} = GenServer.call(server, {:request, init_msg, session_id, %{}}) - - init_notification = build_notification("notifications/initialized", %{}) - - assert :ok = - GenServer.cast( - server, - {:notification, init_notification, session_id, %{}} - ) - - session_name = Anubis.Server.Registry.server_session(StubServer, session_id) - - for _ <- 1..3 do - Process.sleep(100) - ping = build_request("ping", %{}, System.unique_integer()) - assert {:ok, _} = GenServer.call(server, {:request, ping, session_id, %{}}) - assert Session.get(session_name) - end - - Process.sleep(250) - - # After expiration, the session should no longer be accessible - # The session process has been terminated by the supervisor - assert catch_exit(Session.get(session_name)) - end - - test "notifications reset expiry timer" do - transport = start_supervised!(StubTransport) - - server = - start_supervised!( - {Base, - [ - module: StubServer, - name: :notification_reset_server, - transport: [layer: StubTransport, name: transport], - session_idle_timeout: 200 - ]} - ) - - session_id = "notif_session_#{System.unique_integer()}" - - init_msg = - init_request("2025-03-26", %{"name" => "TestClient", "version" => "1.0.0"}) - - assert {:ok, _} = GenServer.call(server, {:request, init_msg, session_id, %{}}) - - init_notification = build_notification("notifications/initialized", %{}) - - assert :ok = - GenServer.cast( - server, - {:notification, init_notification, session_id, %{}} - ) - - session_name = Anubis.Server.Registry.server_session(StubServer, session_id) - - for _ <- 1..3 do - Process.sleep(100) - - notification = - build_notification("notifications/message", %{ - "level" => "info", - "data" => "test" - }) - - assert :ok = - GenServer.cast( - server, - {:notification, notification, session_id, %{}} - ) - - assert Session.get(session_name) - end - - Process.sleep(250) - - # After expiration, the session should no longer be accessible - # The session process has been terminated by the supervisor - assert catch_exit(Session.get(session_name)) - end - end - - describe "sampling requests" do - setup context do - context - |> Map.put(:client_capabilities, %{"sampling" => %{}}) - |> initialized_server() - |> then(fn ctx -> - frame = Frame.put_private(%Frame{}, ctx) - Map.put(ctx, :frame, frame) - end) - end - - test "server can send sampling request to client", %{ - server: server, - transport: transport, - session_id: session_id, - frame: frame - } do - :ok = StubTransport.set_test_pid(transport, self()) - - messages = [ - %{"role" => "user", "content" => %{"type" => "text", "text" => "Hello"}} - ] - - :ok = - Anubis.Server.send_sampling_request(frame, messages, - system_prompt: "You are a helpful assistant", - max_tokens: 100, - metadata: %{test: true} - ) - - Process.sleep(10) - - assert_receive {:send_message, request_data} - assert {:ok, [decoded]} = Message.decode(request_data) - - assert Message.is_request(decoded) - assert decoded["method"] == "sampling/createMessage" - assert decoded["params"]["messages"] == messages - assert decoded["params"]["systemPrompt"] == "You are a helpful assistant" - assert decoded["params"]["maxTokens"] == 100 - - request_id = decoded["id"] - - response = %{ - "id" => request_id, - "result" => %{ - "role" => "assistant", - "content" => %{"type" => "text", "text" => "Hello! How can I help you?"}, - "model" => "test-model", - "stopReason" => "endTurn" - } - } - - :ok = GenServer.cast(server, {:response, response, session_id, %{}}) - - Process.sleep(10) - - state = :sys.get_state(server) - assert state.frame.assigns.last_sampling_response == response["result"] - assert state.frame.assigns.last_sampling_request_id == request_id - end - - test "server handles sampling request timeout", %{ - server: server, - transport: transport, - frame: frame - } do - :ok = StubTransport.set_test_pid(transport, self()) - - messages = [ - %{"role" => "user", "content" => %{"type" => "text", "text" => "Hello"}} - ] - - :ok = Anubis.Server.send_sampling_request(frame, messages) - - Process.sleep(10) - - assert_receive {:send_message, _request_data} - - state = :sys.get_state(server) - assert map_size(state.server_requests) == 1 - end - - test "server handles sampling error response", %{ - server: server, - transport: transport, - session_id: session_id, - frame: frame - } do - :ok = StubTransport.set_test_pid(transport, self()) - - messages = [ - %{"role" => "user", "content" => %{"type" => "text", "text" => "Hello"}} - ] - - :ok = Anubis.Server.send_sampling_request(frame, messages) - - Process.sleep(10) - - assert_receive {:send_message, request_data} - assert {:ok, [decoded]} = Message.decode(request_data) - request_id = decoded["id"] - - error_response = %{ - "id" => request_id, - "error" => %{ - "code" => -32_600, - "message" => "Client doesn't support sampling" - } - } - - :ok = - GenServer.cast(server, {:response, error_response, session_id, %{}}) - - Process.sleep(10) - - state = :sys.get_state(server) - assert map_size(state.server_requests) == 0 - end - end -end diff --git a/test/anubis/server/component/tool_annotations_test.exs b/test/anubis/server/component/tool_annotations_test.exs index 6ff6fbd9..740e2c2c 100644 --- a/test/anubis/server/component/tool_annotations_test.exs +++ b/test/anubis/server/component/tool_annotations_test.exs @@ -2,6 +2,8 @@ defmodule Anubis.Server.Component.ToolAnnotationsTest do use Anubis.MCP.Case, async: true alias Anubis.MCP.Message + alias Anubis.Server.Registry + alias Anubis.Server.Session describe "tool annotations" do test "annotations callback is optional" do @@ -59,46 +61,43 @@ defmodule Anubis.Server.Component.ToolAnnotationsTest do end setup do - start_supervised!(Anubis.Server.Registry) - transport = start_supervised!(StubTransport) + transport_name = Registry.transport_name(ServerWithAnnotatedTools, StubTransport) + start_supervised!({StubTransport, name: transport_name}) - # Start session supervisor - start_supervised!( - {Anubis.Server.Session.Supervisor, server: ServerWithAnnotatedTools, registry: Anubis.Server.Registry} - ) + task_sup = Registry.task_supervisor_name(ServerWithAnnotatedTools) + start_supervised!({Task.Supervisor, name: task_sup}) - server_opts = [ - module: ServerWithAnnotatedTools, - name: :test_server, - registry: Anubis.Server.Registry, - transport: [layer: StubTransport, name: transport] - ] - - server = start_supervised!({Anubis.Server.Base, server_opts}) - - # Initialize the server session_id = "test-session" + session_name = Registry.session_name(ServerWithAnnotatedTools, session_id) + + session = + start_supervised!( + {Session, + session_id: session_id, + server_module: ServerWithAnnotatedTools, + name: session_name, + transport: [layer: StubTransport, name: transport_name], + task_supervisor: task_sup} + ) request = init_request("2025-03-26", %{"name" => "TestClient", "version" => "1.0.0"}) - assert {:ok, _} = GenServer.call(server, {:request, request, session_id, %{}}) + assert {:ok, _} = GenServer.call(session, {:mcp_request, request, %{}}) notification = build_notification("notifications/initialized", %{}) + assert :ok = GenServer.cast(session, {:mcp_notification, notification, %{}}) + Process.sleep(30) - assert :ok = - GenServer.cast(server, {:notification, notification, session_id, %{}}) - - %{server: server, session_id: session_id} + %{server: session, session_id: session_id} end test "lists tools with and without annotations", %{ - server: server, - session_id: session_id + server: server } do request = build_request("tools/list", %{}) {:ok, response_string} = - GenServer.call(server, {:request, request, session_id, %{}}) + GenServer.call(server, {:mcp_request, request, %{}}) {:ok, [response]} = Message.decode(response_string) @@ -132,13 +131,12 @@ defmodule Anubis.Server.Component.ToolAnnotationsTest do end test "lists tools with output schemas", %{ - server: server, - session_id: session_id + server: server } do request = build_request("tools/list", %{}) {:ok, response_string} = - GenServer.call(server, {:request, request, session_id, %{}}) + GenServer.call(server, {:mcp_request, request, %{}}) {:ok, [response]} = Message.decode(response_string) @@ -184,41 +182,39 @@ defmodule Anubis.Server.Component.ToolAnnotationsTest do end setup do - start_supervised!(Anubis.Server.Registry) - transport = start_supervised!(StubTransport) + transport_name = Registry.transport_name(ServerWithOutputSchemaTools, StubTransport) + start_supervised!({StubTransport, name: transport_name}) - # Start session supervisor - start_supervised!( - {Anubis.Server.Session.Supervisor, server: ServerWithOutputSchemaTools, registry: Anubis.Server.Registry} - ) + task_sup = Registry.task_supervisor_name(ServerWithOutputSchemaTools) + start_supervised!({Task.Supervisor, name: task_sup}) - server_opts = [ - module: ServerWithOutputSchemaTools, - name: :test_output_server, - registry: Anubis.Server.Registry, - transport: [layer: StubTransport, name: transport] - ] - - server = start_supervised!({Anubis.Server.Base, server_opts}) - - # Initialize the server session_id = "test-session-output" + session_name = Registry.session_name(ServerWithOutputSchemaTools, session_id) + + session = + start_supervised!( + {Session, + session_id: session_id, + server_module: ServerWithOutputSchemaTools, + name: session_name, + transport: [layer: StubTransport, name: transport_name], + task_supervisor: task_sup}, + id: :output_session + ) request = init_request("2025-03-26", %{"name" => "TestClient", "version" => "1.0.0"}) - assert {:ok, _} = GenServer.call(server, {:request, request, session_id, %{}}) + assert {:ok, _} = GenServer.call(session, {:mcp_request, request, %{}}) notification = build_notification("notifications/initialized", %{}) + assert :ok = GenServer.cast(session, {:mcp_notification, notification, %{}}) + Process.sleep(30) - assert :ok = - GenServer.cast(server, {:notification, notification, session_id, %{}}) - - %{server: server, session_id: session_id} + %{server: session, session_id: session_id} end test "tool with valid output schema returns structured content", %{ - server: server, - session_id: session_id + server: server } do request = build_request("tools/call", %{ @@ -227,7 +223,7 @@ defmodule Anubis.Server.Component.ToolAnnotationsTest do }) {:ok, response_string} = - GenServer.call(server, {:request, request, session_id, %{}}) + GenServer.call(server, {:mcp_request, request, %{}}) {:ok, [response]} = Message.decode(response_string) @@ -253,8 +249,7 @@ defmodule Anubis.Server.Component.ToolAnnotationsTest do end test "tool error response skips output schema validation", %{ - server: server, - session_id: session_id + server: server } do request = build_request("tools/call", %{ @@ -263,14 +258,13 @@ defmodule Anubis.Server.Component.ToolAnnotationsTest do }) {:ok, response_string} = - GenServer.call(server, {:request, request, session_id, %{}}) + GenServer.call(server, {:mcp_request, request, %{}}) assert {:ok, [%{"result" => %{"isError" => true}}]} = Message.decode(response_string) end test "tool without output schema works normally", %{ - server: server, - session_id: session_id + server: server } do request = build_request("tools/call", %{ @@ -279,7 +273,7 @@ defmodule Anubis.Server.Component.ToolAnnotationsTest do }) {:ok, response_string} = - GenServer.call(server, {:request, request, session_id, %{}}) + GenServer.call(server, {:mcp_request, request, %{}}) {:ok, [response]} = Message.decode(response_string) @@ -297,8 +291,7 @@ defmodule Anubis.Server.Component.ToolAnnotationsTest do end test "tool with invalid output fails validation", %{ - server: server, - session_id: session_id + server: server } do request = build_request("tools/call", %{ @@ -307,7 +300,7 @@ defmodule Anubis.Server.Component.ToolAnnotationsTest do }) {:ok, response_string} = - GenServer.call(server, {:request, request, session_id, %{}}) + GenServer.call(server, {:mcp_request, request, %{}}) {:ok, [response]} = Message.decode(response_string) @@ -319,8 +312,7 @@ defmodule Anubis.Server.Component.ToolAnnotationsTest do end test "tool call with missing arguments parameter should not crash server", %{ - server: server, - session_id: session_id + server: server } do request = build_request("tools/call", %{ @@ -328,7 +320,7 @@ defmodule Anubis.Server.Component.ToolAnnotationsTest do }) {:ok, response_string} = - GenServer.call(server, {:request, request, session_id, %{}}) + GenServer.call(server, {:mcp_request, request, %{}}) {:ok, [response]} = Message.decode(response_string) @@ -341,8 +333,7 @@ defmodule Anubis.Server.Component.ToolAnnotationsTest do end test "tool call with missing arguments parameter works for tools without required params", %{ - server: server, - session_id: session_id + server: server } do request = build_request("tools/call", %{ @@ -350,7 +341,7 @@ defmodule Anubis.Server.Component.ToolAnnotationsTest do }) {:ok, response_string} = - GenServer.call(server, {:request, request, session_id, %{}}) + GenServer.call(server, {:mcp_request, request, %{}}) {:ok, [response]} = Message.decode(response_string) diff --git a/test/anubis/server/frame_test.exs b/test/anubis/server/frame_test.exs index 6e943480..06a808e2 100644 --- a/test/anubis/server/frame_test.exs +++ b/test/anubis/server/frame_test.exs @@ -2,8 +2,38 @@ defmodule Anubis.Server.FrameTest do use ExUnit.Case, async: true alias Anubis.Server.Component.Resource + alias Anubis.Server.Context alias Anubis.Server.Frame + describe "assign/2 preserves context" do + test "assigning values does not modify context" do + original_context = %Context{ + session_id: "session-123", + client_info: %{"name" => "test"}, + headers: %{"authorization" => "Bearer token"}, + remote_ip: {127, 0, 0, 1} + } + + frame = %Frame{context: original_context, assigns: %{existing: true}} + updated_frame = Frame.assign(frame, %{new_key: "value", another: 42}) + + assert updated_frame.context == original_context + assert updated_frame.assigns[:new_key] == "value" + assert updated_frame.assigns[:another] == 42 + assert updated_frame.assigns[:existing] == true + end + + test "assigning does not allow overwriting context struct fields" do + context = %Context{session_id: "original"} + frame = %Frame{context: context} + + updated_frame = Frame.assign(frame, %{context: "malicious"}) + + assert updated_frame.context == context + assert updated_frame.assigns[:context] == "malicious" + end + end + describe "register_resource_template/3" do test "registers a resource template at runtime" do frame = Frame.new() diff --git a/test/anubis/server/session/store_test.exs b/test/anubis/server/session/store_test.exs index c42ffbda..a7537af0 100644 --- a/test/anubis/server/session/store_test.exs +++ b/test/anubis/server/session/store_test.exs @@ -1,18 +1,16 @@ defmodule Anubis.Server.Session.StoreTest do use ExUnit.Case, async: false + alias Anubis.Server.Registry alias Anubis.Server.Session - alias Anubis.Server.Session.Supervisor, as: SessionSupervisor alias Anubis.Test.MockSessionStore @moduletag capture_log: true setup do - # Start the mock store start_supervised!(MockSessionStore) MockSessionStore.reset!() - # Configure the application to use the mock store original_config = Application.get_env(:anubis_mcp, :session_store) Application.put_env(:anubis_mcp, :session_store, @@ -33,80 +31,62 @@ defmodule Anubis.Server.Session.StoreTest do end describe "session persistence" do - test "saves session state when initialized" do - session_id = "test_session_123" - start_supervised!({Registry, keys: :unique, name: TestSessionRegistry}) - session_name = {:via, Registry, {TestSessionRegistry, session_id}} - - # Start a session - start_supervised!({Session, session_id: session_id, name: session_name, server_module: TestServer}) - - # Initialize the session - Session.update_from_initialization( - session_name, - "2024-11-21", - %{"name" => "test_client", "version" => "1.0.0"}, - %{"tools" => %{}} - ) + setup do + task_sup = Registry.task_supervisor_name(StubServer) + start_supervised!({Task.Supervisor, name: task_sup}) - Session.mark_initialized(session_name) + transport_name = Registry.transport_name(StubServer, StubTransport) + start_supervised!({StubTransport, name: transport_name}) - # Check that session was persisted - {:ok, stored_state} = MockSessionStore.load(session_id, []) - assert stored_state.protocol_version == "2024-11-21" - assert stored_state.initialized == true - assert stored_state.client_info["name"] == "test_client" + %{transport_name: transport_name, task_sup: task_sup} end - test "restores session from store on startup" do - session_id = "existing_session_456" - start_supervised!({Registry, keys: :unique, name: TestSessionRegistry2}) + test "saves session state when initialized", %{ + transport_name: transport_name, + task_sup: task_sup + } do + session_id = "test_session_123" + session_name = Registry.session_name(StubServer, session_id) + + start_supervised!( + {Session, + session_id: session_id, + server_module: StubServer, + name: session_name, + transport: [layer: StubTransport, name: transport_name], + task_supervisor: task_sup}, + id: :persist_session + ) - # Pre-populate store with session data - session_data = %{ - id: session_id, - protocol_version: "2024-11-21", - initialized: true, - client_info: %{"name" => "restored_client"}, - client_capabilities: %{"tools" => %{}}, - log_level: "info", - pending_requests: %{} + init_request = %{ + "jsonrpc" => "2.0", + "id" => "init_1", + "method" => "initialize", + "params" => %{ + "protocolVersion" => "2025-03-26", + "clientInfo" => %{"name" => "test_client", "version" => "1.0.0"}, + "capabilities" => %{"tools" => %{}} + } } - :ok = MockSessionStore.save(session_id, session_data, []) - - # Start a new session with the same ID - session_name = {:via, Registry, {TestSessionRegistry2, session_id}} + {:ok, _} = GenServer.call(session_name, {:mcp_request, init_request, %{}}) - start_supervised!({Session, session_id: session_id, name: session_name, server_module: TestServer}) - - # Verify the session was restored with the persisted data - session = Session.get(session_name) - assert session.protocol_version == "2024-11-21" - assert session.initialized == true - assert session.client_info["name"] == "restored_client" - end - - test "persists sessions without tokens" do - session_id = "simple_session_789" - start_supervised!({Registry, keys: :unique, name: TestSessionRegistry3}) - session_name = {:via, Registry, {TestSessionRegistry3, session_id}} - - start_supervised!({Session, session_id: session_id, name: session_name, server_module: TestServer}) + init_notif = %{ + "jsonrpc" => "2.0", + "method" => "notifications/initialized" + } - # Mark as initialized to trigger persistence - Session.mark_initialized(session_name) + GenServer.cast(session_name, {:mcp_notification, init_notif, %{}}) + Process.sleep(50) - # Verify session was persisted {:ok, stored_state} = MockSessionStore.load(session_id, []) assert stored_state.initialized == true - assert stored_state.id == session_id + assert stored_state.client_info["name"] == "test_client" end test "handles session updates atomically" do session_id = "update_session_111" - # Save initial session to store initial_data = %{ id: session_id, log_level: "info", @@ -115,7 +95,6 @@ defmodule Anubis.Server.Session.StoreTest do :ok = MockSessionStore.save(session_id, initial_data, []) - # Perform atomic update updates = %{ log_level: "debug", initialized: true @@ -123,7 +102,6 @@ defmodule Anubis.Server.Session.StoreTest do :ok = MockSessionStore.update(session_id, updates, []) - # Verify updates were actually persisted to the store {:ok, stored_session} = MockSessionStore.load(session_id, []) assert stored_session[:log_level] == "debug" assert stored_session[:initialized] == true @@ -131,14 +109,12 @@ defmodule Anubis.Server.Session.StoreTest do end test "lists active sessions" do - # Create multiple sessions session_ids = ["session_a", "session_b", "session_c"] for session_id <- session_ids do MockSessionStore.save(session_id, %{id: session_id}, []) end - # List active sessions {:ok, active} = MockSessionStore.list_active([]) assert length(active) == 3 assert Enum.all?(session_ids, &(&1 in active)) @@ -147,95 +123,62 @@ defmodule Anubis.Server.Session.StoreTest do test "deletes sessions from store" do session_id = "delete_session_222" - # Save a session :ok = MockSessionStore.save(session_id, %{id: session_id}, []) - # Verify it exists assert {:ok, _} = MockSessionStore.load(session_id, []) - # Delete it :ok = MockSessionStore.delete(session_id, []) - # Verify it's gone assert {:error, :not_found} = MockSessionStore.load(session_id, []) end end - describe "session recovery on supervisor startup" do - setup do - # Start a test registry - start_supervised!({Registry, keys: :unique, name: Anubis.Server.Session.StoreTest.TestRegistry}) - :ok - end - - defmodule TestRegistry do - @moduledoc false - alias Anubis.Server.Session.StoreTest.TestRegistry - - def supervisor(:session_supervisor, _server), do: {:via, Registry, {TestRegistry, :supervisor}} - def server_session(_server, session_id), do: {:via, Registry, {TestRegistry, {:session, session_id}}} - - def whereis_server_session(_server, session_id) do - case Registry.lookup(TestRegistry, {:session, session_id}) do - [{pid, _}] -> pid - [] -> nil - end - end - end + describe "session store configuration" do + test "works without store configured" do + Application.delete_env(:anubis_mcp, :session_store) - test "supervisor restores sessions on startup" do - # Pre-populate store with sessions - session_ids = ["restored_1", "restored_2"] + task_sup = Registry.task_supervisor_name(StubServer) + start_supervised!({Task.Supervisor, name: task_sup}) + transport_name = Registry.transport_name(StubServer, StubTransport) + start_supervised!({StubTransport, name: transport_name}) - for session_id <- session_ids do - MockSessionStore.save( - session_id, - %{ - id: session_id, - initialized: true, - protocol_version: "2024-11-21", - log_level: "info" - }, - [] + session_id = "no_store_session" + session_name = Registry.session_name(StubServer, session_id) + + session = + start_supervised!( + {Session, + session_id: session_id, + server_module: StubServer, + name: session_name, + transport: [layer: StubTransport, name: transport_name], + task_supervisor: task_sup}, + id: :no_store_session ) - end - - # Start the supervisor (this should restore sessions) - start_supervised!({SessionSupervisor, server: TestServer, registry: TestRegistry}) - - # Wait a bit for sessions to be restored - Process.sleep(100) - # Verify sessions were restored - for session_id <- session_ids do - pid = TestRegistry.whereis_server_session(TestServer, session_id) - assert is_pid(pid) - - # Get the session and verify it has restored data - session_name = TestRegistry.server_session(TestServer, session_id) - session = Session.get(session_name) - assert session.id == session_id - assert session.initialized == true - end - end - end + init_request = %{ + "jsonrpc" => "2.0", + "id" => "init_1", + "method" => "initialize", + "params" => %{ + "protocolVersion" => "2025-03-26", + "clientInfo" => %{"name" => "test", "version" => "1.0.0"}, + "capabilities" => %{} + } + } - describe "session store configuration" do - test "works without store configured" do - # Remove store configuration - Application.delete_env(:anubis_mcp, :session_store) + {:ok, _} = GenServer.call(session_name, {:mcp_request, init_request, %{}}) - session_id = "no_store_session" - start_supervised!({Registry, keys: :unique, name: TestSessionRegistry5}) - session_name = {:via, Registry, {TestSessionRegistry5, session_id}} + init_notif = %{ + "jsonrpc" => "2.0", + "method" => "notifications/initialized" + } - # Should still be able to create sessions - start_supervised!({Session, session_id: session_id, name: session_name, server_module: TestServer}) + GenServer.cast(session_name, {:mcp_notification, init_notif, %{}}) + Process.sleep(30) - # Session should work normally - Session.mark_initialized(session_name) - session = Session.get(session_name) - assert session.initialized == true + state = :sys.get_state(session) + assert state.initialized == true end end end diff --git a/test/anubis/server/session_test.exs b/test/anubis/server/session_test.exs new file mode 100644 index 00000000..e742244d --- /dev/null +++ b/test/anubis/server/session_test.exs @@ -0,0 +1,362 @@ +defmodule Anubis.Server.SessionTest do + use Anubis.MCP.Case, async: false + + alias Anubis.MCP.Message + alias Anubis.Server.Registry + alias Anubis.Server.Session + + require Message + + @moduletag capture_log: true + + describe "start_link/1" do + test "starts a session with valid options" do + transport_name = Registry.transport_name(StubServer, StubTransport) + start_supervised!({StubTransport, name: transport_name}) + task_sup = Registry.task_supervisor_name(StubServer) + start_supervised!({Task.Supervisor, name: task_sup}) + + session_name = Registry.session_name(StubServer, "test-session") + + assert {:ok, pid} = + Session.start_link( + session_id: "test-session", + server_module: StubServer, + name: session_name, + transport: [layer: StubTransport, name: transport_name], + task_supervisor: task_sup + ) + + assert Process.alive?(pid) + end + end + + describe "handle_call/3 for messages" do + setup :initialized_server + + test "rejects requests when not initialized" do + transport_name = Registry.transport_name(StubServer, StubTransport) + task_sup = Registry.task_supervisor_name(StubServer) + session_name = Registry.session_name(StubServer, "not_initialized") + + session = + start_supervised!( + {Session, + session_id: "not_initialized", + server_module: StubServer, + name: session_name, + transport: [layer: StubTransport, name: transport_name], + task_supervisor: task_sup}, + id: :uninit_session + ) + + request = build_request("tools/list", 123) + + assert {:ok, _} = + GenServer.call(session, {:mcp_request, request, %{}}) + end + + test "accept ping requests when not initialized" do + transport_name = Registry.transport_name(StubServer, StubTransport) + task_sup = Registry.task_supervisor_name(StubServer) + session_name = Registry.session_name(StubServer, "ping_uninit") + + session = + start_supervised!( + {Session, + session_id: "ping_uninit", + server_module: StubServer, + name: session_name, + transport: [layer: StubTransport, name: transport_name], + task_supervisor: task_sup}, + id: :ping_session + ) + + request = build_request("ping", 123) + assert {:ok, _} = GenServer.call(session, {:mcp_request, request, %{}}) + end + end + + describe "handle_cast/2 for notifications" do + setup :initialized_server + + test "handles notifications", %{server: session} do + notification = + build_notification("notifications/cancelled", %{"requestId" => 1}) + + assert :ok = + GenServer.cast(session, {:mcp_notification, notification, %{}}) + end + + test "handles initialize notification", %{server: session} do + notification = build_notification("notifications/initialized", %{}) + + assert :ok = + GenServer.cast(session, {:mcp_notification, notification, %{}}) + end + end + + describe "send_notification/3" do + setup :initialized_server + + test "sends notification to transport", %{server: session} do + assert :ok = + :info + |> Anubis.Server.send_log_message("hello") + |> then(fn _ -> + send( + session, + {:send_notification, "notifications/log/message", %{"level" => :info, "message" => "hello"}} + ) + + :ok + end) + end + end + + describe "session expiration" do + test "session expires after idle timeout" do + transport_name = Registry.transport_name(StubServer, StubTransport) + start_supervised!({StubTransport, name: transport_name}) + task_sup = Registry.task_supervisor_name(StubServer) + start_supervised!({Task.Supervisor, name: task_sup}) + + session_id = "test_session_#{System.unique_integer()}" + session_name = Registry.session_name(StubServer, session_id) + + session = + start_supervised!( + {Session, + session_id: session_id, + server_module: StubServer, + name: session_name, + transport: [layer: StubTransport, name: transport_name], + task_supervisor: task_sup, + session_idle_timeout: 100}, + id: :expiry_session + ) + + init_msg = + init_request("2025-03-26", %{"name" => "TestClient", "version" => "1.0.0"}) + + assert {:ok, _} = GenServer.call(session, {:mcp_request, init_msg, %{}}) + + init_notification = build_notification("notifications/initialized", %{}) + assert :ok = GenServer.cast(session, {:mcp_notification, init_notification, %{}}) + + assert Process.alive?(session) + + Process.sleep(150) + + refute Process.alive?(session) + end + + test "session timer resets on activity" do + transport_name = Registry.transport_name(StubServer, StubTransport) + start_supervised!({StubTransport, name: transport_name}) + task_sup = Registry.task_supervisor_name(StubServer) + start_supervised!({Task.Supervisor, name: task_sup}) + + session_id = "reset_session_#{System.unique_integer()}" + session_name = Registry.session_name(StubServer, session_id) + + session = + start_supervised!( + {Session, + session_id: session_id, + server_module: StubServer, + name: session_name, + transport: [layer: StubTransport, name: transport_name], + task_supervisor: task_sup, + session_idle_timeout: 200}, + id: :reset_session + ) + + init_msg = + init_request("2025-03-26", %{"name" => "TestClient", "version" => "1.0.0"}) + + assert {:ok, _} = GenServer.call(session, {:mcp_request, init_msg, %{}}) + + init_notification = build_notification("notifications/initialized", %{}) + assert :ok = GenServer.cast(session, {:mcp_notification, init_notification, %{}}) + + for _ <- 1..3 do + Process.sleep(100) + ping = build_request("ping", %{}, System.unique_integer()) + assert {:ok, _} = GenServer.call(session, {:mcp_request, ping, %{}}) + assert Process.alive?(session) + end + + Process.sleep(250) + + refute Process.alive?(session) + end + + test "notifications reset expiry timer" do + transport_name = Registry.transport_name(StubServer, StubTransport) + start_supervised!({StubTransport, name: transport_name}) + task_sup = Registry.task_supervisor_name(StubServer) + start_supervised!({Task.Supervisor, name: task_sup}) + + session_id = "notif_session_#{System.unique_integer()}" + session_name = Registry.session_name(StubServer, session_id) + + session = + start_supervised!( + {Session, + session_id: session_id, + server_module: StubServer, + name: session_name, + transport: [layer: StubTransport, name: transport_name], + task_supervisor: task_sup, + session_idle_timeout: 200}, + id: :notif_session + ) + + init_msg = + init_request("2025-03-26", %{"name" => "TestClient", "version" => "1.0.0"}) + + assert {:ok, _} = GenServer.call(session, {:mcp_request, init_msg, %{}}) + + init_notification = build_notification("notifications/initialized", %{}) + assert :ok = GenServer.cast(session, {:mcp_notification, init_notification, %{}}) + + for _ <- 1..3 do + Process.sleep(100) + + notification = + build_notification("notifications/message", %{ + "level" => "info", + "data" => "test" + }) + + assert :ok = + GenServer.cast( + session, + {:mcp_notification, notification, %{}} + ) + + assert Process.alive?(session) + end + + Process.sleep(250) + + refute Process.alive?(session) + end + end + + describe "sampling requests" do + setup context do + context + |> Map.put(:client_capabilities, %{"sampling" => %{}}) + |> initialized_server() + end + + test "server can send sampling request to client", %{ + server: session, + transport: transport + } do + :ok = StubTransport.set_test_pid(transport, self()) + + messages = [ + %{"role" => "user", "content" => %{"type" => "text", "text" => "Hello"}} + ] + + send( + session, + {:send_sampling_request, + %{ + "messages" => messages, + "systemPrompt" => "You are a helpful assistant", + "maxTokens" => 100 + }, 30_000} + ) + + Process.sleep(10) + + assert_receive {:send_message, request_data} + assert {:ok, [decoded]} = Message.decode(request_data) + + assert Message.is_request(decoded) + assert decoded["method"] == "sampling/createMessage" + assert decoded["params"]["messages"] == messages + assert decoded["params"]["systemPrompt"] == "You are a helpful assistant" + assert decoded["params"]["maxTokens"] == 100 + + request_id = decoded["id"] + + response = %{ + "id" => request_id, + "result" => %{ + "role" => "assistant", + "content" => %{"type" => "text", "text" => "Hello! How can I help you?"}, + "model" => "test-model", + "stopReason" => "endTurn" + } + } + + :ok = GenServer.cast(session, {:mcp_response, response, %{}}) + + Process.sleep(10) + + state = :sys.get_state(session) + assert state.frame.assigns.last_sampling_response == response["result"] + assert state.frame.assigns.last_sampling_request_id == request_id + end + + test "server handles sampling request timeout", %{ + server: session, + transport: transport + } do + :ok = StubTransport.set_test_pid(transport, self()) + + messages = [ + %{"role" => "user", "content" => %{"type" => "text", "text" => "Hello"}} + ] + + send(session, {:send_sampling_request, %{"messages" => messages}, 30_000}) + + Process.sleep(10) + + assert_receive {:send_message, _request_data} + + state = :sys.get_state(session) + assert map_size(state.server_requests) == 1 + end + + test "server handles sampling error response", %{ + server: session, + transport: transport + } do + :ok = StubTransport.set_test_pid(transport, self()) + + messages = [ + %{"role" => "user", "content" => %{"type" => "text", "text" => "Hello"}} + ] + + send(session, {:send_sampling_request, %{"messages" => messages}, 30_000}) + + Process.sleep(10) + + assert_receive {:send_message, request_data} + assert {:ok, [decoded]} = Message.decode(request_data) + request_id = decoded["id"] + + error_response = %{ + "id" => request_id, + "error" => %{ + "code" => -32_600, + "message" => "Client doesn't support sampling" + } + } + + :ok = + GenServer.cast(session, {:mcp_response, error_response, %{}}) + + Process.sleep(10) + + state = :sys.get_state(session) + assert map_size(state.server_requests) == 0 + end + end +end diff --git a/test/anubis/server/transport/sse/plug_test.exs b/test/anubis/server/transport/sse/plug_test.exs index 436854db..87ad1e02 100644 --- a/test/anubis/server/transport/sse/plug_test.exs +++ b/test/anubis/server/transport/sse/plug_test.exs @@ -5,15 +5,15 @@ defmodule Anubis.Server.Transport.SSE.PlugTest do import Plug.Conn import Plug.Test + alias Anubis.MCP.Builders alias Anubis.MCP.Message - alias Anubis.Server.Base + alias Anubis.Server.Registry + alias Anubis.Server.Session alias Anubis.Server.Transport.SSE alias Anubis.Server.Transport.SSE.Plug, as: SSEPlug @moduletag capture_log: true - setup :with_default_registry - describe "init/1" do test "requires server option" do assert_raise KeyError, fn -> @@ -33,7 +33,7 @@ defmodule Anubis.Server.Transport.SSE.PlugTest do end end - test "initializes with valid options", %{registry: registry} do + test "initializes with valid options" do opts = SSEPlug.init(server: StubServer, mode: :sse, timeout: 5000) assert %{ @@ -42,28 +42,16 @@ defmodule Anubis.Server.Transport.SSE.PlugTest do timeout: 5000 } = opts - assert transport == registry.transport(StubServer, :sse) - end - - test "uses custom registry when provided" do - start_supervised!(MockCustomRegistry) - assert Process.whereis(MockCustomRegistry) - - opts = - SSEPlug.init(server: StubServer, mode: :sse, registry: MockCustomRegistry) - - expected_transport = MockCustomRegistry.transport(StubServer, :sse) - - assert opts.transport == expected_transport + assert transport == Registry.transport_name(StubServer, :sse) end end describe "SSE endpoint" do - setup %{registry: registry} do - name = registry.transport(StubServer, :sse) + setup do + name = Registry.transport_name(StubServer, :sse) {:ok, transport} = - start_supervised({SSE, server: StubServer, name: name, registry: registry}) + start_supervised({SSE, server: StubServer, name: name}) sse_opts = SSEPlug.init(server: StubServer, mode: :sse) %{sse_opts: sse_opts, transport: transport} @@ -84,7 +72,6 @@ defmodule Anubis.Server.Transport.SSE.PlugTest do endpoint_url = SSE.get_endpoint_url(transport) assert endpoint_url == "/messages" - # Clean up to avoid logs after test ends capture_log(fn -> SSE.unregister_sse_handler(transport, session_id) Process.sleep(10) @@ -114,43 +101,49 @@ defmodule Anubis.Server.Transport.SSE.PlugTest do end describe "POST endpoint" do - setup %{registry: registry} do - # Start the session supervisor - {:ok, _session_sup} = - start_supervised({ - Anubis.Server.Session.Supervisor, - server: StubServer, registry: registry - }) + setup do + task_sup = Registry.task_supervisor_name(StubServer) + start_supervised!({Task.Supervisor, name: task_sup}) - # Start a stub transport for the server - stub_transport = - start_supervised!({StubTransport, name: registry.transport(StubServer, :stub)}) - - # Start the Base server with stub transport - {:ok, _server} = - start_supervised({ - Base, - module: StubServer, - name: registry.server(StubServer), - transport: [layer: StubTransport, name: stub_transport], - registry: registry - }) + transport_name = Registry.transport_name(StubServer, StubTransport) + start_supervised!({StubTransport, name: transport_name}) - # Now start the SSE transport - name = registry.transport(StubServer, :sse) + session_id = "test-session" + session_name = Registry.session_name(StubServer, session_id) + + start_supervised!( + {Session, + session_id: session_id, + server_module: StubServer, + name: session_name, + transport: [layer: StubTransport, name: transport_name], + task_supervisor: task_sup}, + id: :sse_post_session + ) + + init_request = + Builders.init_request(nil, %{"name" => "Test", "version" => "1.0"}, %{}) + + {:ok, _} = GenServer.call(session_name, {:mcp_request, init_request, %{}}) + + init_notif = Builders.build_notification("notifications/initialized", %{}) + GenServer.cast(session_name, {:mcp_notification, init_notif, %{}}) + Process.sleep(30) + + name = Registry.transport_name(StubServer, :sse) {:ok, transport} = - start_supervised({SSE, server: StubServer, name: name, registry: registry}) + start_supervised({SSE, server: StubServer, name: name}) post_opts = SSEPlug.init(server: StubServer, mode: :post) - %{post_opts: post_opts, transport: transport, registry: registry} + %{post_opts: post_opts, transport: transport, session_id: session_id} end test "POST request with valid JSON returns response", %{ post_opts: post_opts, - transport: transport + transport: transport, + session_id: session_id } do - session_id = "test-session" :ok = SSE.register_sse_handler(transport, session_id) request = build_request("ping", %{}) @@ -166,14 +159,13 @@ defmodule Anubis.Server.Transport.SSE.PlugTest do assert conn.status == 202 assert conn.resp_body == "{}" - # Clean up to avoid logs after test ends capture_log(fn -> SSE.unregister_sse_handler(transport, session_id) Process.sleep(10) end) end - test "POST request with notification returns 202", %{post_opts: post_opts} do + test "POST request with notification returns 202", %{post_opts: post_opts, session_id: session_id} do notification = build_notification("notifications/message", %{ "level" => "info", @@ -186,6 +178,7 @@ defmodule Anubis.Server.Transport.SSE.PlugTest do :post |> conn("/messages", body) |> put_req_header("content-type", "application/json") + |> put_req_header("x-session-id", session_id) |> SSEPlug.call(post_opts) assert conn.status == 202 @@ -218,17 +211,17 @@ defmodule Anubis.Server.Transport.SSE.PlugTest do end describe "session ID extraction" do - setup %{registry: registry} do - name = registry.transport(StubServer, :sse) + setup do + name = Registry.transport_name(StubServer, :sse) {:ok, transport} = - start_supervised({SSE, server: StubServer, name: name, registry: registry}) + start_supervised({SSE, server: StubServer, name: name}) post_opts = SSEPlug.init(server: StubServer, mode: :post) %{post_opts: post_opts, transport: transport} end - test "extracts session ID from header", %{post_opts: post_opts} do + test "notification to unknown session is accepted", %{post_opts: post_opts} do notification = build_notification("notifications/message", %{ "level" => "info", @@ -240,30 +233,14 @@ defmodule Anubis.Server.Transport.SSE.PlugTest do conn = :post |> conn("/messages", body) - |> put_req_header("x-session-id", "header-session-123") - |> SSEPlug.call(post_opts) - - assert conn.status == 202 - end - - test "extracts session ID from query params", %{post_opts: post_opts} do - notification = - build_notification("notifications/message", %{ - "level" => "info", - "data" => "test" - }) - - {:ok, body} = Message.encode_notification(notification) - - conn = - :post - |> conn("/messages?session_id=query-session-456", body) + |> put_req_header("x-session-id", "nonexistent-session") |> SSEPlug.call(post_opts) + # SSE transport accepts notifications (fire-and-forget) assert conn.status == 202 end - test "generates session ID if not provided", %{post_opts: post_opts} do + test "notification without session ID is accepted", %{post_opts: post_opts} do notification = build_notification("notifications/message", %{ "level" => "info", diff --git a/test/anubis/server/transport/sse_test.exs b/test/anubis/server/transport/sse_test.exs index ca7aaafd..f76b5151 100644 --- a/test/anubis/server/transport/sse_test.exs +++ b/test/anubis/server/transport/sse_test.exs @@ -3,10 +3,10 @@ defmodule Anubis.Server.Transport.SSETest do import ExUnit.CaptureLog + alias Anubis.Server.Registry + alias Anubis.Server.Session alias Anubis.Server.Transport.SSE - setup :with_default_registry - describe "start_link/1" do test "starts with valid options" do server = :"test_server_#{System.unique_integer([:positive])}" @@ -46,11 +46,10 @@ defmodule Anubis.Server.Transport.SSETest do describe "with running transport" do setup do - registry = Anubis.Server.Registry - name = registry.transport(StubServer, :sse) + name = Registry.transport_name(StubServer, :sse) {:ok, transport} = - start_supervised({SSE, server: StubServer, name: name, registry: registry}) + start_supervised({SSE, server: StubServer, name: name}) %{transport: transport, server: StubServer} end @@ -68,13 +67,30 @@ defmodule Anubis.Server.Transport.SSETest do test "handle_message processes notifications", %{transport: transport} do session_id = "test-session-456" + task_sup = Registry.task_supervisor_name(StubServer) + start_supervised!({Task.Supervisor, name: task_sup}) + + transport_name = Registry.transport_name(StubServer, StubTransport) + start_supervised!({StubTransport, name: transport_name}, id: :sse_stub_transport) + + session_name = Registry.session_name(StubServer, session_id) + + start_supervised!( + {Session, + session_id: session_id, + server_module: StubServer, + name: session_name, + transport: [layer: StubTransport, name: transport_name], + task_supervisor: task_sup}, + id: :sse_session + ) + notification = build_notification("notifications/message", %{ "level" => "info", "data" => "test" }) - # Should return nil for notifications and send them to server assert {:ok, nil} = SSE.handle_message(transport, session_id, notification, %{}) end @@ -89,7 +105,6 @@ defmodule Anubis.Server.Transport.SSETest do assert_receive {:sse_message, ^message} - # Clean up to avoid logs after test ends capture_log(fn -> SSE.unregister_sse_handler(transport, session_id) Process.sleep(10) @@ -108,10 +123,8 @@ defmodule Anubis.Server.Transport.SSETest do session1 = "session-1" session2 = "session-2" - # Register two handlers assert :ok = SSE.register_sse_handler(transport, session1) - # Second handler in a different process test_pid = self() spawn(fn -> @@ -128,11 +141,9 @@ defmodule Anubis.Server.Transport.SSETest do message = "broadcast message" assert :ok = SSE.send_message(transport, message, timeout: 5000) - # Both handlers should receive the message assert_receive {:sse_message, ^message} assert_receive {:handler2_received, ^message} - # Clean up to avoid logs after test ends capture_log(fn -> SSE.unregister_sse_handler(transport, session1) SSE.unregister_sse_handler(transport, session2) @@ -172,17 +183,11 @@ defmodule Anubis.Server.Transport.SSETest do end test "get_endpoint_url with custom base_url and post_path" do - registry = Anubis.Server.Registry name = :custom_sse_transport {:ok, transport} = start_supervised( - {SSE, - server: StubServer, - name: name, - base_url: "http://localhost:8080", - post_path: "/api/messages", - registry: registry}, + {SSE, server: StubServer, name: name, base_url: "http://localhost:8080", post_path: "/api/messages"}, id: :custom_sse ) @@ -192,13 +197,11 @@ defmodule Anubis.Server.Transport.SSETest do test "shutdown/1 gracefully shuts down", %{transport: transport} do session_id = "shutdown-test" - # Register a handler assert :ok = SSE.register_sse_handler(transport, session_id) assert Process.alive?(transport) assert :ok = SSE.shutdown(transport) - # Should send close message to handler assert_receive :close_sse Process.sleep(100) diff --git a/test/anubis/server/transport/streamable_http/plug_persistence_test.exs b/test/anubis/server/transport/streamable_http/plug_persistence_test.exs index 124f0038..474c6717 100644 --- a/test/anubis/server/transport/streamable_http/plug_persistence_test.exs +++ b/test/anubis/server/transport/streamable_http/plug_persistence_test.exs @@ -4,15 +4,14 @@ defmodule Anubis.Server.Transport.StreamableHTTP.PlugPersistenceTest do import Plug.Conn import Plug.Test + alias Anubis.Server.Supervisor, as: ServerSupervisor alias Anubis.Server.Transport.StreamableHTTP.Plug alias Anubis.Test.MockSessionStore setup do - # Start the mock store {:ok, _} = MockSessionStore.start_link([]) MockSessionStore.reset!() - # Configure the application to use the mock store original_config = Application.get_env(:anubis_mcp, :session_store) Application.put_env(:anubis_mcp, :session_store, @@ -34,21 +33,28 @@ defmodule Anubis.Server.Transport.StreamableHTTP.PlugPersistenceTest do describe "session persistence without tokens" do setup do - # We need a minimal test server setup - opts = [ - server: TestServer, - registry: Anubis.Server.Registry, - session_header: "mcp-session-id" - ] - - plug_opts = Plug.init(opts) + session_config = %{ + server_module: StubServer, + registry_mod: Anubis.Server.Registry.None, + transport: [layer: StubTransport, name: :stub_transport], + session_idle_timeout: nil, + timeout: 30_000, + task_supervisor: :test_task_sup + } + + :persistent_term.put({ServerSupervisor, StubServer, :session_config}, session_config) + + on_exit(fn -> + :persistent_term.erase({ServerSupervisor, StubServer, :session_config}) + end) + + plug_opts = Plug.init(server: StubServer, session_header: "mcp-session-id") {:ok, plug_opts: plug_opts} end test "GET request with existing session ID reconnects to stored session", %{plug_opts: _plug_opts} do session_id = "existing_session_123" - # Pre-populate store with session session_data = %{ id: session_id, protocol_version: "2024-11-21", @@ -58,38 +64,31 @@ defmodule Anubis.Server.Transport.StreamableHTTP.PlugPersistenceTest do :ok = MockSessionStore.save(session_id, session_data, []) - # Make GET request with session ID conn = :get |> conn("/") |> put_req_header("accept", "text/event-stream") |> put_req_header("mcp-session-id", session_id) - # The plug should accept the reconnection based on session ID alone - # (backward compatible behavior - no token required) assert get_req_header(conn, "mcp-session-id") == [session_id] - # Verify session still exists in store assert {:ok, stored_data} = MockSessionStore.load(session_id, []) assert stored_data.id == session_id assert stored_data.initialized == true end test "GET request without session ID generates new session", %{plug_opts: _plug_opts} do - # Make GET request without session ID conn = :get |> conn("/") |> put_req_header("accept", "text/event-stream") - # Should work fine without a session ID (new session will be created) assert get_req_header(conn, "mcp-session-id") == [] end test "POST request with existing session ID uses stored session", %{plug_opts: _plug_opts} do session_id = "post_session_111" - # Pre-populate store with session session_data = %{ id: session_id, protocol_version: "2024-11-21", @@ -98,7 +97,6 @@ defmodule Anubis.Server.Transport.StreamableHTTP.PlugPersistenceTest do :ok = MockSessionStore.save(session_id, session_data, []) - # Make POST request with session ID message = %{ "jsonrpc" => "2.0", "method" => "tools/list", @@ -112,19 +110,40 @@ defmodule Anubis.Server.Transport.StreamableHTTP.PlugPersistenceTest do |> put_req_header("accept", "application/json, text/event-stream") |> put_req_header("mcp-session-id", session_id) - # Should accept the request based on session ID alone assert get_req_header(conn, "mcp-session-id") == [session_id] end end describe "session lifecycle with persistence" do - test "initialize request triggers session persistence" do - # Make initialize request - message = %{ + setup do + task_sup = :"test_task_sup_#{System.unique_integer([:positive])}" + transport_name = :"test_transport_#{System.unique_integer([:positive])}" + + start_supervised!({Task.Supervisor, name: task_sup}) + start_supervised!({StubTransport, name: transport_name}) + + %{task_supervisor: task_sup, transport_name: transport_name} + end + + test "initialize request triggers session persistence", ctx do + session_id = "persist_init_#{System.unique_integer([:positive])}" + session_name = :"session_#{session_id}" + + {:ok, session} = + start_supervised( + {Anubis.Server.Session, + session_id: session_id, + server_module: StubServer, + name: session_name, + transport: [layer: StubTransport, name: ctx.transport_name], + task_supervisor: ctx.task_supervisor} + ) + + init_request = %{ "jsonrpc" => "2.0", "method" => "initialize", "params" => %{ - "protocolVersion" => "2024-11-21", + "protocolVersion" => "2025-03-26", "capabilities" => %{}, "clientInfo" => %{ "name" => "test_client", @@ -134,22 +153,17 @@ defmodule Anubis.Server.Transport.StreamableHTTP.PlugPersistenceTest do "id" => "init_1" } - _conn = - :post - |> conn("/", JSON.encode!(message)) - |> put_req_header("content-type", "application/json") - |> put_req_header("accept", "application/json") + {:ok, _response} = GenServer.call(session, {:mcp_request, init_request, %{}}) - # This test just validates the request structure - # In a real scenario, the transport would handle the initialization - # and persist the session to the store - assert true + {:ok, stored} = MockSessionStore.load(session_id, []) + assert stored.id == session_id + assert stored.protocol_version == "2025-03-26" + assert stored.client_info == %{"name" => "test_client", "version" => "1.0.0"} end test "session data can be updated in store" do session_id = "update_session_222" - # Save initial session initial_data = %{ id: session_id, initialized: false, @@ -158,7 +172,6 @@ defmodule Anubis.Server.Transport.StreamableHTTP.PlugPersistenceTest do :ok = MockSessionStore.save(session_id, initial_data, []) - # Update session updates = %{ initialized: true, log_level: "debug" @@ -166,7 +179,6 @@ defmodule Anubis.Server.Transport.StreamableHTTP.PlugPersistenceTest do :ok = MockSessionStore.update(session_id, updates, []) - # Verify updates were applied {:ok, updated_data} = MockSessionStore.load(session_id, []) assert updated_data.initialized == true assert updated_data.log_level == "debug" @@ -176,7 +188,6 @@ defmodule Anubis.Server.Transport.StreamableHTTP.PlugPersistenceTest do test "DELETE request removes session from store" do session_id = "delete_session_333" - # Pre-populate store with session session_data = %{ id: session_id, initialized: true @@ -184,32 +195,26 @@ defmodule Anubis.Server.Transport.StreamableHTTP.PlugPersistenceTest do :ok = MockSessionStore.save(session_id, session_data, []) - # Verify session exists assert {:ok, _} = MockSessionStore.load(session_id, []) - # Make DELETE request conn = :delete |> conn("/") |> put_req_header("mcp-session-id", session_id) - # Simulate session deletion (would be handled by transport) :ok = MockSessionStore.delete(session_id, []) - # Session should be removed from store assert {:error, :not_found} = MockSessionStore.load(session_id, []) assert get_req_header(conn, "mcp-session-id") == [session_id] end test "can list all active sessions" do - # Create multiple sessions session_ids = ["session_a", "session_b", "session_c"] for session_id <- session_ids do :ok = MockSessionStore.save(session_id, %{id: session_id}, []) end - # List active sessions {:ok, active} = MockSessionStore.list_active([]) assert length(active) == 3 assert Enum.all?(session_ids, &(&1 in active)) @@ -218,10 +223,8 @@ defmodule Anubis.Server.Transport.StreamableHTTP.PlugPersistenceTest do describe "backward compatibility" do test "sessions work exactly as before when no store is configured" do - # Remove store configuration Application.delete_env(:anubis_mcp, :session_store) - # Make a request - should work fine without persistence message = %{ "jsonrpc" => "2.0", "method" => "ping", @@ -234,26 +237,22 @@ defmodule Anubis.Server.Transport.StreamableHTTP.PlugPersistenceTest do |> put_req_header("content-type", "application/json") |> put_req_header("accept", "application/json") - # Request should work normally even without store assert conn.request_path == "/" assert get_req_header(conn, "content-type") == ["application/json"] end test "clients without session IDs work normally" do - # Client doesn't send session ID (first connection) conn = :get |> conn("/") |> put_req_header("accept", "text/event-stream") - # Should work fine - server will generate session ID assert conn.request_path == "/" end test "clients with session IDs reconnect transparently" do session_id = "client_session_999" - # Save session to simulate previous connection :ok = MockSessionStore.save( session_id, @@ -265,14 +264,12 @@ defmodule Anubis.Server.Transport.StreamableHTTP.PlugPersistenceTest do [] ) - # Client reconnects with same session ID conn = :get |> conn("/") |> put_req_header("accept", "text/event-stream") |> put_req_header("mcp-session-id", session_id) - # Should reconnect to existing session transparently assert get_req_header(conn, "mcp-session-id") == [session_id] end end diff --git a/test/anubis/server/transport/streamable_http/plug_test.exs b/test/anubis/server/transport/streamable_http/plug_test.exs index b1577c09..ee669640 100644 --- a/test/anubis/server/transport/streamable_http/plug_test.exs +++ b/test/anubis/server/transport/streamable_http/plug_test.exs @@ -6,19 +6,46 @@ defmodule Anubis.Server.Transport.StreamableHTTP.PlugTest do import Plug.Test alias Anubis.MCP.Message + alias Anubis.Server.Registry + alias Anubis.Server.Supervisor, as: ServerSupervisor alias Anubis.Server.Transport.StreamableHTTP alias Anubis.Server.Transport.StreamableHTTP.Plug, as: StreamableHTTPPlug - setup :with_default_registry + defp setup_session_config(opts \\ []) do + task_sup = Registry.task_supervisor_name(StubServer) + transport_name = Registry.transport_name(StubServer, StubTransport) + + session_config = %{ + server_module: StubServer, + registry_mod: Keyword.get(opts, :registry_mod, Registry.None), + transport: [layer: StubTransport, name: transport_name], + session_idle_timeout: nil, + timeout: 30_000, + task_supervisor: task_sup + } + + :persistent_term.put({ServerSupervisor, StubServer, :session_config}, session_config) + session_config + end + + defp cleanup_session_config do + :persistent_term.erase({ServerSupervisor, StubServer, :session_config}) + end describe "init/1" do + setup do + setup_session_config() + on_exit(&cleanup_session_config/0) + :ok + end + test "requires server option" do assert_raise KeyError, fn -> StreamableHTTPPlug.init([]) end end - test "initializes with valid options", %{registry: registry} do + test "initializes with valid options" do opts = StreamableHTTPPlug.init(server: StubServer) assert %{ @@ -27,10 +54,10 @@ defmodule Anubis.Server.Transport.StreamableHTTP.PlugTest do timeout: 30_000 } = opts - assert transport == registry.transport(StubServer, :streamable_http) + assert transport == Registry.transport_name(StubServer, :streamable_http) end - test "accepts custom session header", %{registry: registry} do + test "accepts custom session header" do opts = StreamableHTTPPlug.init( server: StubServer, @@ -43,33 +70,20 @@ defmodule Anubis.Server.Transport.StreamableHTTP.PlugTest do timeout: 30_000 } = opts - assert transport == registry.transport(StubServer, :streamable_http) - end - - test "uses custom registry when provided" do - start_supervised!(MockCustomRegistry) - assert Process.whereis(MockCustomRegistry) - - opts = - StreamableHTTPPlug.init( - server: StubServer, - mode: :streamable_http, - registry: MockCustomRegistry - ) - - expected_transport = MockCustomRegistry.transport(StubServer, :streamable_http) - - assert opts.transport == expected_transport + assert transport == Registry.transport_name(StubServer, :streamable_http) end end describe "GET endpoint" do - setup %{registry: registry} do - name = registry.transport(StubServer, :streamable_http) - sup = registry.task_supervisor(StubServer) + setup do + setup_session_config() + on_exit(&cleanup_session_config/0) + + name = Registry.transport_name(StubServer, :streamable_http) + sup = Registry.task_supervisor_name(StubServer) {:ok, transport} = - start_supervised({StreamableHTTP, server: StubServer, name: name, registry: registry, task_supervisor: sup}) + start_supervised({StreamableHTTP, server: StubServer, name: name, task_supervisor: sup}) opts = StreamableHTTPPlug.init(server: StubServer) @@ -88,10 +102,6 @@ defmodule Anubis.Server.Transport.StreamableHTTP.PlugTest do session_id = "test-session-123" assert :ok = StreamableHTTP.register_sse_handler(transport, session_id) - # Note: We don't actually call the plug here because it would - # establish a persistent connection and hang the test - - # Clean up to avoid logs after test ends capture_log(fn -> StreamableHTTP.unregister_sse_handler(transport, session_id) Process.sleep(10) @@ -111,42 +121,69 @@ defmodule Anubis.Server.Transport.StreamableHTTP.PlugTest do end describe "POST endpoint" do - setup %{registry: registry} do - # Start the session supervisor - {:ok, _session_sup} = - start_supervised({ - Anubis.Server.Session.Supervisor, - server: StubServer, registry: registry - }) + setup do + task_sup = Registry.task_supervisor_name(StubServer) + start_supervised!({Task.Supervisor, name: task_sup}) - # Start a stub transport for the server - stub_transport = - start_supervised!({StubTransport, name: registry.transport(StubServer, :stub)}) - - # Start the Base server with stub transport - {:ok, _server} = - start_supervised({ - Anubis.Server.Base, - module: StubServer, - name: registry.server(StubServer), - transport: [layer: StubTransport, name: stub_transport], - registry: registry - }) + transport_name = Registry.transport_name(StubServer, StubTransport) + start_supervised!({StubTransport, name: transport_name}) + + registry_name = Registry.registry_name(StubServer) + start_supervised!({Registry.Local, name: registry_name}) - # Now start the StreamableHTTP transport - name = registry.transport(StubServer, :streamable_http) - sup = registry.task_supervisor(StubServer) - start_supervised!({Task.Supervisor, name: sup}) + session_config = setup_session_config(registry_mod: Registry.Local) + on_exit(&cleanup_session_config/0) + + session_sup_name = Registry.session_supervisor_name(StubServer) + start_supervised!({DynamicSupervisor, name: session_sup_name, strategy: :one_for_one}) + + name = Registry.transport_name(StubServer, :streamable_http) {:ok, transport} = - start_supervised({StreamableHTTP, server: StubServer, name: name, registry: registry, task_supervisor: sup}) + start_supervised({StreamableHTTP, server: StubServer, name: name, task_supervisor: task_sup}) opts = StreamableHTTPPlug.init(server: StubServer) - %{opts: opts, transport: transport} + test_session_id = "post-test-session" + session_name = Registry.session_name(StubServer, test_session_id) + + {:ok, _session} = + ServerSupervisor.start_session(StubServer, + session_id: test_session_id, + server_module: StubServer, + name: session_name, + transport: session_config.transport, + session_idle_timeout: 1_800_000, + timeout: 30_000, + task_supervisor: task_sup + ) + + Registry.Local.register_session(registry_name, test_session_id, Process.whereis(session_name)) + + init_req = %{ + "jsonrpc" => "2.0", + "id" => "setup_init", + "method" => "initialize", + "params" => %{ + "protocolVersion" => "2025-03-26", + "clientInfo" => %{"name" => "Test", "version" => "1.0"}, + "capabilities" => %{} + } + } + + {:ok, _} = GenServer.call(session_name, {:mcp_request, init_req, %{}}) + + GenServer.cast( + session_name, + {:mcp_notification, %{"jsonrpc" => "2.0", "method" => "notifications/initialized"}, %{}} + ) + + Process.sleep(30) + + %{opts: opts, transport: transport, test_session_id: test_session_id} end - test "POST request with notification returns 202", %{opts: opts} do + test "POST request with notification returns 202", %{opts: opts, test_session_id: session_id} do notification = build_notification("notifications/message", %{ "level" => "info", @@ -159,14 +196,15 @@ defmodule Anubis.Server.Transport.StreamableHTTP.PlugTest do :post |> conn("/", body) |> put_req_header("content-type", "application/json") - |> put_req_header("accept", "application/json, text/event-stream") + |> put_req_header("accept", "application/json") + |> put_req_header("mcp-session-id", session_id) |> StreamableHTTPPlug.call(opts) assert conn.status == 202 assert conn.resp_body == "{}" end - test "POST request with valid request returns response", %{opts: opts} do + test "POST request with valid request returns response", %{opts: opts, test_session_id: session_id} do request = build_request("ping", %{}) {:ok, body} = Message.encode_request(request, 1) @@ -174,7 +212,8 @@ defmodule Anubis.Server.Transport.StreamableHTTP.PlugTest do :post |> conn("/", body) |> put_req_header("content-type", "application/json") - |> put_req_header("accept", "application/json, text/event-stream") + |> put_req_header("accept", "application/json") + |> put_req_header("mcp-session-id", session_id) |> StreamableHTTPPlug.call(opts) assert conn.status == 200 @@ -187,7 +226,7 @@ defmodule Anubis.Server.Transport.StreamableHTTP.PlugTest do :post |> conn("/", "invalid json") |> put_req_header("content-type", "application/json") - |> put_req_header("accept", "application/json, text/event-stream") + |> put_req_header("accept", "application/json") |> StreamableHTTPPlug.call(opts) assert conn.status == 400 @@ -197,12 +236,15 @@ defmodule Anubis.Server.Transport.StreamableHTTP.PlugTest do end describe "DELETE endpoint" do - setup %{registry: registry} do - name = registry.transport(StubServer, :streamable_http) - sup = registry.task_supervisor(StubServer) + setup do + setup_session_config() + on_exit(&cleanup_session_config/0) + + name = Registry.transport_name(StubServer, :streamable_http) + sup = Registry.task_supervisor_name(StubServer) {:ok, transport} = - start_supervised({StreamableHTTP, server: StubServer, name: name, registry: registry, task_supervisor: sup}) + start_supervised({StreamableHTTP, server: StubServer, name: name, task_supervisor: sup}) opts = StreamableHTTPPlug.init(server: StubServer) @@ -233,12 +275,15 @@ defmodule Anubis.Server.Transport.StreamableHTTP.PlugTest do end describe "unsupported methods" do - setup %{registry: registry} do - name = registry.transport(StubServer, :streamable_http) - sup = registry.task_supervisor(StubServer) + setup do + setup_session_config() + on_exit(&cleanup_session_config/0) + + name = Registry.transport_name(StubServer, :streamable_http) + sup = Registry.task_supervisor_name(StubServer) {:ok, _transport} = - start_supervised({StreamableHTTP, server: StubServer, name: name, registry: registry, task_supervisor: sup}) + start_supervised({StreamableHTTP, server: StubServer, name: name, task_supervisor: sup}) opts = StreamableHTTPPlug.init(server: StubServer) @@ -258,42 +303,69 @@ defmodule Anubis.Server.Transport.StreamableHTTP.PlugTest do end describe "session handling" do - setup %{registry: registry} do - # Start the session supervisor - {:ok, _session_sup} = - start_supervised({ - Anubis.Server.Session.Supervisor, - server: StubServer, registry: registry - }) + setup do + task_sup = Registry.task_supervisor_name(StubServer) + start_supervised!({Task.Supervisor, name: task_sup}) - # Start a stub transport for the server - stub_transport = - start_supervised!({StubTransport, name: registry.transport(StubServer, :stub)}) - - # Start the Base server with stub transport - {:ok, _server} = - start_supervised({ - Anubis.Server.Base, - module: StubServer, - name: registry.server(StubServer), - transport: [layer: StubTransport, name: stub_transport], - registry: registry - }) + transport_name = Registry.transport_name(StubServer, StubTransport) + start_supervised!({StubTransport, name: transport_name}) + + registry_name = Registry.registry_name(StubServer) + start_supervised!({Registry.Local, name: registry_name}) + + session_config = setup_session_config(registry_mod: Registry.Local) + on_exit(&cleanup_session_config/0) + + session_sup_name = Registry.session_supervisor_name(StubServer) + start_supervised!({DynamicSupervisor, name: session_sup_name, strategy: :one_for_one}) - # Now start the StreamableHTTP transport - name = registry.transport(StubServer, :streamable_http) - sup = registry.task_supervisor(StubServer) - start_supervised!({Task.Supervisor, name: sup}) + name = Registry.transport_name(StubServer, :streamable_http) {:ok, transport} = - start_supervised({StreamableHTTP, server: StubServer, name: name, registry: registry, task_supervisor: sup}) + start_supervised({StreamableHTTP, server: StubServer, name: name, task_supervisor: task_sup}) opts = StreamableHTTPPlug.init(server: StubServer) - %{opts: opts, transport: transport} + test_session_id = "session-handling-test" + session_name = Registry.session_name(StubServer, test_session_id) + + {:ok, _session} = + ServerSupervisor.start_session(StubServer, + session_id: test_session_id, + server_module: StubServer, + name: session_name, + transport: session_config.transport, + session_idle_timeout: 1_800_000, + timeout: 30_000, + task_supervisor: task_sup + ) + + Registry.Local.register_session(registry_name, test_session_id, Process.whereis(session_name)) + + init_req = %{ + "jsonrpc" => "2.0", + "id" => "setup_init", + "method" => "initialize", + "params" => %{ + "protocolVersion" => "2025-03-26", + "clientInfo" => %{"name" => "Test", "version" => "1.0"}, + "capabilities" => %{} + } + } + + {:ok, _} = GenServer.call(session_name, {:mcp_request, init_req, %{}}) + + GenServer.cast( + session_name, + {:mcp_notification, %{"jsonrpc" => "2.0", "method" => "notifications/initialized"}, %{}} + ) + + Process.sleep(30) + + %{opts: opts, transport: transport, test_session_id: test_session_id} end - test "extracts session ID from header", %{opts: opts} do + test "extracts session ID from header", %{opts: opts, test_session_id: session_id} do notification = build_notification("notifications/message", %{ "level" => "info", @@ -306,14 +378,14 @@ defmodule Anubis.Server.Transport.StreamableHTTP.PlugTest do :post |> conn("/", body) |> put_req_header("content-type", "application/json") - |> put_req_header("accept", "application/json, text/event-stream") - |> put_req_header("mcp-session-id", "header-session-123") + |> put_req_header("accept", "application/json") + |> put_req_header("mcp-session-id", session_id) |> StreamableHTTPPlug.call(opts) assert conn.status == 202 end - test "generates session ID if not provided", %{opts: opts} do + test "notification to unknown session returns 400", %{opts: opts} do notification = build_notification("notifications/message", %{ "level" => "info", @@ -326,17 +398,19 @@ defmodule Anubis.Server.Transport.StreamableHTTP.PlugTest do :post |> conn("/", body) |> put_req_header("content-type", "application/json") - |> put_req_header("accept", "application/json, text/event-stream") + |> put_req_header("accept", "application/json") + |> put_req_header("mcp-session-id", "unknown-session") |> StreamableHTTPPlug.call(opts) - assert conn.status == 202 + assert conn.status == 400 end - test "initialize request generates new session ID", %{opts: opts} do + test "initialize request creates new session", %{opts: opts} do init_request = build_request("initialize", %{ "protocolVersion" => "2025-03-26", - "clientInfo" => %{"name" => "test", "version" => "1.0.0"} + "clientInfo" => %{"name" => "test", "version" => "1.0.0"}, + "capabilities" => %{} }) {:ok, body} = Message.encode_request(init_request, 1) @@ -345,11 +419,12 @@ defmodule Anubis.Server.Transport.StreamableHTTP.PlugTest do :post |> conn("/", body) |> put_req_header("content-type", "application/json") - |> put_req_header("accept", "application/json, text/event-stream") - |> put_req_header("mcp-session-id", "should-be-ignored") + |> put_req_header("accept", "application/json") |> StreamableHTTPPlug.call(opts) assert conn.status == 200 + {:ok, response} = Jason.decode(conn.resp_body) + assert response["result"]["protocolVersion"] end end end diff --git a/test/anubis/server/transport/streamable_http_test.exs b/test/anubis/server/transport/streamable_http_test.exs index 891dbd1e..fe18a44d 100644 --- a/test/anubis/server/transport/streamable_http_test.exs +++ b/test/anubis/server/transport/streamable_http_test.exs @@ -3,24 +3,19 @@ defmodule Anubis.Server.Transport.StreamableHTTPTest do import ExUnit.CaptureLog + alias Anubis.Server.Registry alias Anubis.Server.Transport.StreamableHTTP - alias Anubis.Server.Transport.StreamableHTTP.RequestParams - - setup :with_default_registry describe "start_link/1" do test "starts with valid options" do - server = Anubis.Server.Registry.server(StubServer) - name = Anubis.Server.Registry.transport(StubServer, :streamable_http) - sup = Anubis.Server.Registry.task_supervisor(StubServer) + server = :"test_server_#{System.unique_integer([:positive])}" + name = Registry.transport_name(server, :streamable_http) + sup = Registry.task_supervisor_name(server) assert {:ok, pid} = StreamableHTTP.start_link(server: server, name: name, task_supervisor: sup) assert Process.alive?(pid) - - assert Anubis.Server.Registry.whereis_transport(StubServer, :streamable_http) == - pid end test "requires server option" do @@ -32,13 +27,12 @@ defmodule Anubis.Server.Transport.StreamableHTTPTest do describe "with running transport" do setup do - registry = Anubis.Server.Registry - name = registry.transport(StubServer, :streamable_http) - sup = registry.task_supervisor(StubServer) + name = Registry.transport_name(StubServer, :streamable_http) + sup = Registry.task_supervisor_name(StubServer) start_supervised!({Task.Supervisor, name: sup}) {:ok, transport} = - start_supervised({StreamableHTTP, server: StubServer, name: name, registry: registry, task_supervisor: sup}) + start_supervised({StreamableHTTP, server: StubServer, name: name, task_supervisor: sup}) %{transport: transport, server: StubServer} end @@ -53,32 +47,6 @@ defmodule Anubis.Server.Transport.StreamableHTTPTest do refute StreamableHTTP.get_sse_handler(transport, session_id) end - test "handle_message_for_sse fails when server is not in registry", %{ - transport: transport - } do - session_id = "test-session-456" - - assert :ok = StreamableHTTP.register_sse_handler(transport, session_id) - message = build_request("ping", %{}) - - params = %RequestParams{ - transport: transport, - session_id: session_id, - message: message, - context: %{}, - session_header: nil, - timeout: 5 - } - - StreamableHTTP.handle_message_for_sse(params) - - # Clean up to avoid logs after test ends - capture_log(fn -> - StreamableHTTP.unregister_sse_handler(transport, session_id) - Process.sleep(10) - end) - end - test "routes messages to sessions", %{transport: transport} do session_id = "test-session-789" @@ -89,7 +57,6 @@ defmodule Anubis.Server.Transport.StreamableHTTPTest do assert_receive {:sse_message, ^message} - # Clean up to avoid logs after test ends capture_log(fn -> StreamableHTTP.unregister_sse_handler(transport, session_id) Process.sleep(10) diff --git a/test/support/mcp/assertions.ex b/test/support/mcp/assertions.ex index 7a9bbc2d..fc91634a 100644 --- a/test/support/mcp/assertions.ex +++ b/test/support/mcp/assertions.ex @@ -1,7 +1,7 @@ defmodule Anubis.MCP.Assertions do @moduledoc false - import ExUnit.Assertions, only: [assert: 2, assert: 1] + import ExUnit.Assertions, only: [assert: 2] def assert_client_initialized(client) when is_pid(client) do state = :sys.get_state(client) @@ -10,12 +10,6 @@ defmodule Anubis.MCP.Assertions do def assert_server_initialized(server) when is_pid(server) do state = :sys.get_state(server) - assert {session_id, _} = state.sessions |> Map.to_list() |> List.first() - - assert session = - Anubis.Server.Registry.whereis_server_session(StubServer, session_id) - - state = :sys.get_state(session) - assert state.initialized, "Expected server to be initialized" + assert state.initialized, "Expected server session to be initialized" end end diff --git a/test/support/mcp/setup.ex b/test/support/mcp/setup.ex index 6ff1e939..d67045df 100644 --- a/test/support/mcp/setup.ex +++ b/test/support/mcp/setup.ex @@ -7,9 +7,9 @@ defmodule Anubis.MCP.Setup do alias Anubis.MCP.Builders alias Anubis.MCP.Message - alias Anubis.Server.Base + alias Anubis.Server.Registry alias Anubis.Server.Session - alias Anubis.Server.Transport + alias Anubis.Server.Transport.STDIO require Message @@ -72,169 +72,79 @@ defmodule Anubis.MCP.Setup do Process.sleep(50) end - def initialized_client_with_server(ctx) do - protocol_version = ctx[:protocol_version] - capabilities = ctx[:client_capabilities] - info = ctx[:client_info] || %{"name" => "TestClient", "version" => "1.0.0"} - - start_supervised!(Anubis.Server.Registry) - transport = start_supervised!(StubTransport) - - client_opts = [ - transport: [layer: StubTransport, name: transport], - client_info: info, - capabilities: capabilities, - protocol_version: protocol_version - ] - - client = start_supervised!({Anubis.Client.Base, client_opts}) - unique_id = System.unique_integer([:positive]) - start_supervised!({StubServer, transport: StubTransport}, id: unique_id) - assert server = Anubis.Server.Registry.whereis_server(StubServer) - - Process.sleep(30) - - StubTransport.set_client(transport, client) - - Process.sleep(80) - - assert_client_initialized(client) - assert_server_initialized(server) - - :ok = StubTransport.clear(transport) - - Map.merge(ctx, %{transport: transport, client: client, server: server}) - end - def initialized_server(ctx) do session_id = ctx[:session_id] || "test-session-123" protocol_version = ctx[:protocol_version] capabilities = ctx[:client_capabilities] info = ctx[:client_info] || %{"name" => "TestClient", "version" => "1.0.0"} - start_supervised!(Anubis.Server.Registry) - transport = start_supervised!(StubTransport) + transport_name = Registry.transport_name(StubServer, StubTransport) + transport = start_supervised!({StubTransport, name: transport_name}) - # Start session supervisor - start_supervised!({Anubis.Server.Session.Supervisor, server: StubServer, registry: Anubis.Server.Registry}) + task_sup = Registry.task_supervisor_name(StubServer) + start_supervised!({Task.Supervisor, name: task_sup}) - server_opts = [ - module: StubServer, - name: Anubis.Server.Registry.server(StubServer), - registry: Anubis.Server.Registry, - transport: [layer: StubTransport, name: transport] - ] + session_name = Registry.session_name(StubServer, session_id) - server = start_supervised!({Base, server_opts}) - assert server == Anubis.Server.Registry.whereis_server(StubServer) + session = + start_supervised!( + {Session, + session_id: session_id, + server_module: StubServer, + name: session_name, + transport: [layer: StubTransport, name: transport_name], + task_supervisor: task_sup} + ) request = Builders.init_request(protocol_version, info, capabilities) - assert {:ok, _} = GenServer.call(server, {:request, request, session_id, %{}}) + assert {:ok, _} = GenServer.call(session, {:mcp_request, request, %{}}) notification = Builders.build_notification("notifications/initialized", %{}) - - assert :ok = - GenServer.cast(server, {:notification, notification, session_id, %{}}) + assert :ok = GenServer.cast(session, {:mcp_notification, notification, %{}}) Process.sleep(50) - assert_server_initialized(server) + assert_server_initialized(session) :ok = StubTransport.clear(transport) Map.merge(ctx, %{ transport: transport, - server: server, + server: session, session_id: session_id, - server_registry: Anubis.Server.Registry, server_module: StubServer }) end - def initialized_base_server(ctx) do - server_module = StubServer - session_id = ctx[:session_id] || "test-session-123" - protocol_version = ctx[:protocol_version] - capabilities = ctx[:client_capabilities] - transport = ctx[:transport] || StubTransport - info = ctx[:client_info] || %{"name" => "TestClient", "version" => "1.0.0"} - - # session supervisor - %{registry: registry} = ctx = with_default_registry(ctx) - - start_supervised!({Session.Supervisor, server: server_module, registry: ctx.registry}) - - assert registry.supervisor(server_module, :session_supervisor) - - # base server - server_name = registry.server(server_module) - transport_name = registry.transport(server_module, transport) - - server_opts = [ - module: server_module, - name: server_name, - registry: registry, - transport: [ - layer: transport, - name: transport_name - ] - ] - - start_supervised!({Base, server_opts}) - assert server = registry.whereis_server(server_module) - - # transport - start_supervised!({transport, name: transport_name, server: server_name, registry: registry}) - - assert registry.whereis_transport(server_module, transport) - - request = Builders.init_request(protocol_version, info, capabilities) - assert {:ok, _} = GenServer.call(server, {:request, request, session_id}) - notification = Builders.build_notification("notifications/initialized", %{}) - assert :ok = GenServer.cast(server, {:notification, notification, session_id}) - - Process.sleep(50) - - assert_server_initialized(server) - - :ok = StubTransport.clear(transport) - - Map.merge(ctx, %{transport: transport, server: server, session_id: session_id}) - end - def server_with_stdio_transport(ctx) do - name = ctx[:name] || :test_stdio_server - name = Anubis.Server.Registry.server(name) server_module = ctx[:server_module] || StubServer + transport_name = Registry.transport_name(server_module, :stdio) + task_sup = Registry.task_supervisor_name(server_module) + start_supervised!({Task.Supervisor, name: task_sup}) - transport_name = Anubis.Server.Registry.transport(server_module, :stdio) - start_supervised!({Transport.STDIO, name: transport_name, server: server_module}) - - assert transport = - Anubis.Server.Registry.whereis_transport(server_module, :stdio) + session_name = Registry.stdio_session_name(server_module) - opts = [ - module: server_module, - name: name, - transport: [layer: Transport.STDIO, name: transport_name] - ] + session = + start_supervised!( + {Session, + session_id: "stdio", + server_module: server_module, + name: session_name, + transport: [ + layer: STDIO, + name: transport_name + ], + task_supervisor: task_sup} + ) - start_supervised!({Base, opts}) - assert server = Anubis.Server.Registry.whereis_server(server_module) + transport = + start_supervised!({STDIO, name: transport_name, server: server_module}) - Map.merge(ctx, %{server: server, transport: transport}) - end - - def with_default_registry(ctx) do - start_supervised!(Anubis.Server.Registry) - assert Process.whereis(Anubis.Server.Registry) - Map.put(ctx, :registry, Anubis.Server.Registry) + Map.merge(ctx, %{server: session, transport: transport}) end def initialized_client(context) do import Mox - start_supervised!(Anubis.Server.Registry) - server_capabilities = context[:server_capabilities] || %{ diff --git a/test/support/mock_custom_registry.ex b/test/support/mock_custom_registry.ex index 1fcab890..9a88fa60 100644 --- a/test/support/mock_custom_registry.ex +++ b/test/support/mock_custom_registry.ex @@ -1,82 +1,59 @@ defmodule MockCustomRegistry do @moduledoc false - @behaviour Anubis.Server.Registry.Adapter + @behaviour Anubis.Server.Registry - alias Anubis.Server.Registry.Adapter + use GenServer - @impl true + @impl Anubis.Server.Registry def child_spec(opts) do %{ id: __MODULE__, - start: {GenServer, :start_link, [__MODULE__, opts, [name: __MODULE__]]}, + start: {__MODULE__, :start_link, [opts]}, type: :worker, restart: :permanent, shutdown: 500 } end - @impl Adapter - def transport(server, transport_type) do - {:via, __MODULE__, {:transport, server, transport_type}} + @impl Anubis.Server.Registry + def register_session(name, session_id, pid) do + GenServer.call(name, {:register, session_id, pid}) end - @impl Adapter - def task_supervisor(server_module) do - {:via, __MODULE__, {:task_supervisor, server_module}} + @impl Anubis.Server.Registry + def lookup_session(name, session_id) do + GenServer.call(name, {:lookup, session_id}) end - @impl Adapter - def server(server_module) do - {:via, __MODULE__, {:server, server_module}} + @impl Anubis.Server.Registry + def unregister_session(name, session_id) do + GenServer.call(name, {:unregister, session_id}) end - @impl Adapter - def server_session(server_module, session_id) do - {:via, __MODULE__, {:server_session, server_module, session_id}} - end - - @impl Adapter - def supervisor(kind, server_module) do - {:via, __MODULE__, {:supervisor, kind, server_module}} - end - - @impl Adapter - def whereis_server(server_module) do - case :ets.lookup(__MODULE__, {:server, server_module}) do - [{_, pid}] -> pid - [] -> nil - end + def start_link(opts \\ []) do + GenServer.start_link(__MODULE__, opts, name: __MODULE__) end - @impl Adapter - def whereis_server_session(server_module, session_id) do - case :ets.lookup(__MODULE__, {:server_session, server_module, session_id}) do - [{_, pid}] -> pid - [] -> nil - end + @impl GenServer + def init(_opts) do + {:ok, %{sessions: %{}}} end - @impl Adapter - def whereis_transport(server_module, transport_type) do - case :ets.lookup(__MODULE__, {:transport, server_module, transport_type}) do - [{_, pid}] -> pid - [] -> nil - end + @impl GenServer + def handle_call({:register, session_id, pid}, _from, state) do + sessions = Map.put(state.sessions, session_id, pid) + {:reply, :ok, %{state | sessions: sessions}} end - @impl Adapter - def whereis_supervisor(kind, server_module) do - case :ets.lookup(__MODULE__, {:supervisor, kind, server_module}) do - [{_, pid}] -> pid - [] -> nil + def handle_call({:lookup, session_id}, _from, state) do + case Map.get(state.sessions, session_id) do + nil -> {:reply, {:error, :not_found}, state} + pid -> {:reply, {:ok, pid}, state} end end - def start_link(opts \\ []) do - GenServer.start_link(__MODULE__, opts, name: __MODULE__) - end - - def init(_opts) do - {:ok, %{}} + def handle_call({:unregister, session_id}, _from, state) do + sessions = Map.delete(state.sessions, session_id) + {:reply, :ok, %{state | sessions: sessions}} end end diff --git a/test/support/stub_transport.ex b/test/support/stub_transport.ex index 158c99ef..623a0874 100644 --- a/test/support/stub_transport.ex +++ b/test/support/stub_transport.ex @@ -29,7 +29,7 @@ defmodule StubTransport do """ @impl true def start_link(opts \\ []) do - state = %{messages: [], client: nil, server: nil, test_pid: nil} + state = %{messages: [], client: nil, test_pid: nil} if name = opts[:name] do GenServer.start_link(__MODULE__, state, name: name) @@ -127,16 +127,15 @@ defmodule StubTransport do def handle_call({:send_message, message}, _from, state) do new_messages = [message | state.messages] - # Send to test process if configured if state.test_pid do send(state.test_pid, {:send_message, message}) end if is_binary(message) do message = decode_message(message) - forward_to_server(message, state) + forward_to_session(message, state) else - forward_to_server(message, state) + forward_to_session(message, state) end {:reply, :ok, %{state | messages: new_messages}} @@ -156,26 +155,33 @@ defmodule StubTransport do message end - defp forward_to_server(message, state) when Message.is_request(message) do + defp forward_to_session(message, state) when Message.is_request(message) do if message["method"] == "sampling/createMessage" do :ok else - name = Anubis.Server.Registry.server(StubServer) + session_name = + Anubis.Server.Registry.session_name(StubServer, state.session_id) {:ok, response} = - GenServer.call(name, {:request, message, state.session_id, %{}}) + GenServer.call(session_name, {:mcp_request, message, %{}}) - GenServer.cast(state.client, {:response, response}) + if state.client do + GenServer.cast(state.client, {:response, response}) + end end end - defp forward_to_server(message, state) when Message.is_response(message) or Message.is_error(message) do - name = Anubis.Server.Registry.server(StubServer) - GenServer.cast(name, {:response, message, state.session_id, %{}}) + defp forward_to_session(message, state) when Message.is_response(message) or Message.is_error(message) do + session_name = + Anubis.Server.Registry.session_name(StubServer, state.session_id) + + GenServer.cast(session_name, {:mcp_response, message, %{}}) end - defp forward_to_server(message, state) when Message.is_notification(message) do - name = Anubis.Server.Registry.server(StubServer) - :ok = GenServer.cast(name, {:notification, message, state.session_id}) + defp forward_to_session(message, state) when Message.is_notification(message) do + session_name = + Anubis.Server.Registry.session_name(StubServer, state.session_id) + + :ok = GenServer.cast(session_name, {:mcp_notification, message, %{}}) end end