Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
52 changes: 29 additions & 23 deletions lib/anubis/client/base.ex
Original file line number Diff line number Diff line change
Expand Up @@ -112,6 +112,8 @@ defmodule Anubis.Client.Base do
optional(:sampling | String.t()) => %{}
}

@default_operation_timeout to_timeout(second: 30)

@typedoc """
MCP client initialization options

Expand All @@ -136,7 +138,8 @@ defmodule Anubis.Client.Base do
{:transport, {:required, {:custom, &Anubis.client_transport/1}}},
{:client_info, {:required, :map}},
{:capabilities, {:required, :map}},
{:protocol_version, {:string, {:default, @default_protocol_version}}}
{:protocol_version, {:string, {:default, @default_protocol_version}}},
{:timeout, {:integer, {:default, @default_operation_timeout}}}
])

@doc """
Expand Down Expand Up @@ -172,7 +175,7 @@ defmodule Anubis.Client.Base do
method: "ping",
params: %{},
progress_opts: Keyword.get(opts, :progress),
timeout: Keyword.get(opts, :timeout)
timeout: Keyword.get(opts, :timeout, @default_operation_timeout)
})

buffer_timeout = operation.timeout + to_timeout(second: 1)
Expand Down Expand Up @@ -200,7 +203,7 @@ defmodule Anubis.Client.Base do
method: "resources/list",
params: params,
progress_opts: Keyword.get(opts, :progress),
timeout: Keyword.get(opts, :timeout)
timeout: Keyword.get(opts, :timeout, @default_operation_timeout)
})

buffer_timeout = operation.timeout + to_timeout(second: 1)
Expand Down Expand Up @@ -228,7 +231,7 @@ defmodule Anubis.Client.Base do
method: "resources/templates/list",
params: params,
progress_opts: Keyword.get(opts, :progress),
timeout: Keyword.get(opts, :timeout)
timeout: Keyword.get(opts, :timeout, @default_operation_timeout)
})

buffer_timeout = operation.timeout + to_timeout(second: 1)
Expand All @@ -253,7 +256,7 @@ defmodule Anubis.Client.Base do
method: "resources/read",
params: %{"uri" => uri},
progress_opts: Keyword.get(opts, :progress),
timeout: Keyword.get(opts, :timeout)
timeout: Keyword.get(opts, :timeout, @default_operation_timeout)
})

buffer_timeout = operation.timeout + to_timeout(second: 1)
Expand Down Expand Up @@ -281,7 +284,7 @@ defmodule Anubis.Client.Base do
method: "prompts/list",
params: params,
progress_opts: Keyword.get(opts, :progress),
timeout: Keyword.get(opts, :timeout)
timeout: Keyword.get(opts, :timeout, @default_operation_timeout)
})

buffer_timeout = operation.timeout + to_timeout(second: 1)
Expand Down Expand Up @@ -309,7 +312,7 @@ defmodule Anubis.Client.Base do
method: "prompts/get",
params: params,
progress_opts: Keyword.get(opts, :progress),
timeout: Keyword.get(opts, :timeout)
timeout: Keyword.get(opts, :timeout, @default_operation_timeout)
})

buffer_timeout = operation.timeout + to_timeout(second: 1)
Expand Down Expand Up @@ -337,7 +340,7 @@ defmodule Anubis.Client.Base do
method: "tools/list",
params: params,
progress_opts: Keyword.get(opts, :progress),
timeout: Keyword.get(opts, :timeout)
timeout: Keyword.get(opts, :timeout, @default_operation_timeout)
})

buffer_timeout = operation.timeout + to_timeout(second: 1)
Expand Down Expand Up @@ -365,7 +368,7 @@ defmodule Anubis.Client.Base do
method: "tools/call",
params: params,
progress_opts: Keyword.get(opts, :progress),
timeout: Keyword.get(opts, :timeout)
timeout: Keyword.get(opts, :timeout, @default_operation_timeout)
})

buffer_timeout = operation.timeout + to_timeout(second: 1)
Expand Down Expand Up @@ -418,7 +421,8 @@ defmodule Anubis.Client.Base do
operation =
Operation.new(%{
method: "logging/setLevel",
params: %{"level" => level}
params: %{"level" => level},
timeout: @default_operation_timeout
})

buffer_timeout = operation.timeout + to_timeout(second: 1)
Expand Down Expand Up @@ -474,7 +478,7 @@ defmodule Anubis.Client.Base do
method: "completion/complete",
params: params,
progress_opts: Keyword.get(opts, :progress),
timeout: Keyword.get(opts, :timeout)
timeout: Keyword.get(opts, :timeout, @default_operation_timeout)
})

buffer_timeout = operation.timeout + to_timeout(second: 1)
Expand Down Expand Up @@ -785,7 +789,8 @@ defmodule Anubis.Client.Base do
client_info: opts.client_info,
capabilities: opts.capabilities,
protocol_version: protocol_version,
transport: transport
transport: transport,
timeout: opts.timeout
})

client_name = get_in(opts, [:client_info, "name"])
Expand Down Expand Up @@ -827,7 +832,7 @@ defmodule Anubis.Client.Base do
{request_id, updated_state} =
State.add_request_from_operation(state, operation, from),
{:ok, request_data} <- encode_request(method, params_with_token, request_id),
:ok <- send_to_transport(state.transport, request_data) do
:ok <- send_to_transport(state.transport, request_data, timeout: operation.timeout) do
Telemetry.execute(
Telemetry.event_client_request(),
%{system_time: System.system_time()},
Expand Down Expand Up @@ -885,7 +890,7 @@ defmodule Anubis.Client.Base do
"progress" => progress,
"total" => total
}) do
send_to_transport(state.transport, notification)
send_to_transport(state.transport, notification, timeout: state.timeout)
end, state}
end

Expand Down Expand Up @@ -972,14 +977,15 @@ defmodule Anubis.Client.Base do
operation =
Operation.new(%{
method: "initialize",
params: params
params: params,
timeout: state.timeout
})

{request_id, updated_state} =
State.add_request_from_operation(state, operation, {self(), make_ref()})

with {:ok, request_data} <- encode_request("initialize", params, request_id),
:ok <- send_to_transport(state.transport, request_data) do
:ok <- send_to_transport(state.transport, request_data, timeout: operation.timeout) do
{:noreply, updated_state}
else
err -> {:stop, err, state}
Expand Down Expand Up @@ -1018,7 +1024,7 @@ defmodule Anubis.Client.Base do

with {:ok, response_data} <-
Message.encode_response(%{"result" => roots_result}, id),
:ok <- send_to_transport(state.transport, response_data) do
:ok <- send_to_transport(state.transport, response_data, timeout: state.timeout) do
Logging.client_event("roots_list_request", %{id: id, roots_count: roots_count})

Telemetry.execute(
Expand All @@ -1044,7 +1050,7 @@ defmodule Anubis.Client.Base do

defp handle_server_request(%{"method" => "ping", "id" => id}, state) do
with {:ok, response_data} <- Message.encode_response(%{"result" => %{}}, id),
:ok <- send_to_transport(state.transport, response_data) do
:ok <- send_to_transport(state.transport, response_data, timeout: state.timeout) do
{:noreply, state}
else
err ->
Expand Down Expand Up @@ -1494,15 +1500,15 @@ defmodule Anubis.Client.Base do
send_notification(state, "notifications/cancelled", params)
end

defp send_to_transport(transport, data) do
with {:error, reason} <- transport.layer.send_message(transport.name, data) do
defp send_to_transport(transport, data, opts) do
with {:error, reason} <- transport.layer.send_message(transport.name, data, opts) do
{:error, Error.transport(:send_failure, %{original_reason: reason})}
end
end

defp send_notification(state, method, params \\ %{}) do
with {:ok, notification_data} <- encode_notification(method, params) do
send_to_transport(state.transport, notification_data)
send_to_transport(state.transport, notification_data, timeout: state.timeout)
end
end

Expand Down Expand Up @@ -1593,7 +1599,7 @@ defmodule Anubis.Client.Base do

defp send_sampling_response(id, response, state) do
transport = state.transport
:ok = transport.layer.send_message(transport.name, response)
:ok = transport.layer.send_message(transport.name, response, timeout: state.timeout)

Telemetry.execute(
Telemetry.event_client_response(),
Expand All @@ -1605,7 +1611,7 @@ defmodule Anubis.Client.Base do
defp send_sampling_error(id, message, code, reason, %{transport: transport} = state) do
error = %Error{code: -1, message: message, data: %{"reason" => reason}}
{:ok, response} = Error.to_json_rpc(error, id)
:ok = transport.layer.send_message(transport.name, response)
:ok = transport.layer.send_message(transport.name, response, timeout: state.timeout)

Logging.client_event(
"sampling_error",
Expand Down
10 changes: 4 additions & 6 deletions lib/anubis/client/operation.ex
Original file line number Diff line number Diff line change
Expand Up @@ -9,8 +9,6 @@ defmodule Anubis.Client.Operation do
- `timeout` - The timeout for this specific operation (default: 30 seconds)
"""

@default_timeout to_timeout(second: 30)

@type progress_options :: [
token: String.t() | integer(),
callback: (String.t() | integer(), number(), number() | nil -> any())
Expand All @@ -25,9 +23,9 @@ defmodule Anubis.Client.Operation do

defstruct [
:method,
:timeout,
params: %{},
progress_opts: [],
timeout: @default_timeout
progress_opts: []
]

@doc """
Expand All @@ -47,12 +45,12 @@ defmodule Anubis.Client.Operation do
optional(:progress_opts) => progress_options() | nil,
optional(:timeout) => pos_integer()
}) :: t()
def new(%{method: method} = attrs) do
def new(%{method: method, timeout: timeout} = attrs) do
%__MODULE__{
method: method,
params: Map.get(attrs, :params) || %{},
progress_opts: Map.get(attrs, :progress_opts),
timeout: Map.get(attrs, :timeout) || @default_timeout
timeout: timeout
}
end
end
5 changes: 4 additions & 1 deletion lib/anubis/client/state.ex
Original file line number Diff line number Diff line change
Expand Up @@ -14,6 +14,7 @@ defmodule Anubis.Client.State do
server_capabilities: map() | nil,
server_info: map() | nil,
protocol_version: String.t(),
timeout: pos_integer(),
transport: map(),
pending_requests: %{String.t() => Request.t()},
progress_callbacks: %{String.t() => Base.progress_callback()},
Expand All @@ -28,6 +29,7 @@ defmodule Anubis.Client.State do
:capabilities,
:server_capabilities,
:server_info,
:timeout,
:protocol_version,
:transport,
pending_requests: %{},
Expand All @@ -43,7 +45,8 @@ defmodule Anubis.Client.State do
client_info: opts.client_info,
capabilities: opts.capabilities,
protocol_version: opts.protocol_version,
transport: opts.transport
transport: opts.transport,
timeout: opts.timeout
}
end

Expand Down
20 changes: 11 additions & 9 deletions lib/anubis/server/base.ex
Original file line number Diff line number Diff line change
Expand Up @@ -60,7 +60,8 @@ defmodule Anubis.Server.Base do
{: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}}}
{: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()
Expand Down Expand Up @@ -90,7 +91,8 @@ defmodule Anubis.Server.Base do
session_idle_timeout: opts.session_idle_timeout,
expiry_timers: %{},
frame: Frame.new(),
server_requests: %{}
server_requests: %{},
timeout: opts.timeout
}

Logging.server_event("starting", %{
Expand Down Expand Up @@ -233,7 +235,7 @@ defmodule Anubis.Server.Base do

def handle_info({:send_notification, method, params}, state) do
with {:ok, notification} <- encode_notification(method, params),
:ok <- send_to_transport(state.transport, notification) do
:ok <- send_to_transport(state.transport, notification, timeout: state.timeout) do
{:noreply, state}
else
{:error, err} ->
Expand Down Expand Up @@ -674,12 +676,12 @@ defmodule Anubis.Server.Base do
Message.encode_notification(notification)
end

defp send_to_transport(nil, _data) do
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) do
with {:error, reason} <- layer.send_message(name, data) do
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
Expand Down Expand Up @@ -723,7 +725,7 @@ defmodule Anubis.Server.Base do
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) do
: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
Expand Down Expand Up @@ -854,7 +856,7 @@ defmodule Anubis.Server.Base do

with :ok <- validate_client_capability(state, "roots"),
{:ok, request_data} <- encode_request("roots/list", %{}, request_id),
:ok <- send_to_transport(state.transport, request_data) do
: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
Expand Down Expand Up @@ -890,7 +892,7 @@ defmodule Anubis.Server.Base do
"requestId" => request_id,
"reason" => "timeout"
}),
:ok <- send_to_transport(state.transport, notification) do
:ok <- send_to_transport(state.transport, notification, timeout: state.timeout) do
Logging.server_event(
"roots_request_timeout_cancelled",
%{request_id: request_id}
Expand Down
5 changes: 2 additions & 3 deletions lib/anubis/server/transport/sse.ex
Original file line number Diff line number Diff line change
Expand Up @@ -131,9 +131,8 @@ defmodule Anubis.Server.Transport.SSE do
* `{:error, reason}` otherwise
"""
@impl Transport
@spec send_message(GenServer.server(), binary()) :: :ok | {:error, term()}
def send_message(transport, message) when is_binary(message) do
GenServer.call(transport, {:send_message, message})
def send_message(transport, message, opts) when is_binary(message) do
GenServer.call(transport, {:send_message, message}, opts[:timeout])
end

@doc """
Expand Down
9 changes: 5 additions & 4 deletions lib/anubis/server/transport/stdio.ex
Original file line number Diff line number Diff line change
Expand Up @@ -77,9 +77,8 @@ defmodule Anubis.Server.Transport.STDIO do
* `{:error, reason}` otherwise
"""
@impl Transport
@spec send_message(GenServer.server(), binary()) :: :ok | {:error, term()}
def send_message(transport, message) when is_binary(message) do
GenServer.cast(transport, {:send, message})
def send_message(transport, message, opts) when is_binary(message) do
GenServer.call(transport, {:send, message}, opts[:timeout])
end
Comment thread
zoedsoupe marked this conversation as resolved.

@doc """
Expand Down Expand Up @@ -262,7 +261,9 @@ defmodule Anubis.Server.Transport.STDIO do
else
case GenServer.call(server, {:request, message, "stdio", context}, timeout) do
{:ok, response} when is_binary(response) ->
send_message(self(), response)
# send_message(self(), response)
# NOTE: will be fixed soon, we need to rewrite stdio for server
:ok
Comment thread
zoedsoupe marked this conversation as resolved.

{:error, reason} ->
Logging.transport_event("server_error", %{reason: reason}, level: :error)
Expand Down
5 changes: 2 additions & 3 deletions lib/anubis/server/transport/streamable_http.ex
Original file line number Diff line number Diff line change
Expand Up @@ -118,9 +118,8 @@ defmodule Anubis.Server.Transport.StreamableHTTP do
* `{:error, reason}` otherwise
"""
@impl Transport
@spec send_message(GenServer.server(), binary()) :: :ok | {:error, term()}
def send_message(transport, message) when is_binary(message) do
GenServer.call(transport, {:send_message, message}, 5000)
def send_message(transport, message, opts) when is_binary(message) do
GenServer.call(transport, {:send_message, message}, opts[:timeout])
end
Comment thread
zoedsoupe marked this conversation as resolved.

@doc """
Expand Down
Loading