import asyncio
import json
from typing import Any, Callable, Dict, List, Optional, Union
from dataclasses import dataclass
from enum import Enum
from .transport import Transport, TcpTransport, StdioTransport, SseTransport
from .errors import DcpError, ConnectionError, TimeoutError, ProtocolError
class ProtocolVersion(Enum):
V2024_11_05 = "2024-11-05"
V2025_03_26 = "2025-03-26"
V2025_06_18 = "2025-06-18"
@classmethod
def from_str(cls, version: str) -> Optional["ProtocolVersion"]:
for v in cls:
if v.value == version:
return v
return None
def supports_roots(self) -> bool:
return self in (ProtocolVersion.V2025_03_26, ProtocolVersion.V2025_06_18)
def supports_elicitation(self) -> bool:
return self == ProtocolVersion.V2025_06_18
@dataclass
class Root:
uri: str
name: Optional[str] = None
def to_dict(self) -> Dict[str, Any]:
result = {"uri": self.uri}
if self.name is not None:
result["name"] = self.name
return result
@classmethod
def from_dict(cls, data: Dict[str, Any]) -> "Root":
return cls(uri=data["uri"], name=data.get("name"))
@dataclass
class ElicitationRequest:
message: str
requested_schema: Optional[Dict[str, Any]] = None
def to_dict(self) -> Dict[str, Any]:
result = {"message": self.message}
if self.requested_schema is not None:
result["requestedSchema"] = self.requested_schema
return result
@dataclass
class ElicitationResponse:
action: str content: Optional[Dict[str, Any]] = None
@classmethod
def from_dict(cls, data: Dict[str, Any]) -> "ElicitationResponse":
return cls(action=data["action"], content=data.get("content"))
@dataclass
class ResourceTemplate:
uri_template: str
name: str
description: Optional[str] = None
mime_type: Optional[str] = None
@classmethod
def from_dict(cls, data: Dict[str, Any]) -> "ResourceTemplate":
return cls(
uri_template=data["uriTemplate"],
name=data["name"],
description=data.get("description"),
mime_type=data.get("mimeType")
)
def substitute(self, params: Dict[str, str]) -> str:
result = self.uri_template
for key, value in params.items():
result = result.replace(f"{{{key}}}", value)
return result
class DcpClient:
DEFAULT_PROTOCOL_VERSION = ProtocolVersion.V2025_06_18
DCP_CAPABILITY_KEYS = ("tools", "resources", "prompts", "extensions")
def __init__(self, transport: Transport, timeout: float = 30.0,
protocol_version: Optional[ProtocolVersion] = None,
dcp_capabilities: Optional[Dict[str, List[int]]] = None):
self.transport = transport
self.timeout = timeout
self._preferred_version = protocol_version or self.DEFAULT_PROTOCOL_VERSION
self._dcp_capabilities = self._normalize_dcp_capabilities(dcp_capabilities)
self._negotiated_version: Optional[ProtocolVersion] = None
self._request_id = 0
self._pending: Dict[int, asyncio.Future] = {}
self._notification_handlers: Dict[str, Callable] = {}
self._receive_task: Optional[asyncio.Task] = None
self._initialized = False
self._server_capabilities: Dict[str, Any] = {}
self._roots_changed_callback: Optional[Callable[[List[Root]], Any]] = None
@classmethod
def _normalize_dcp_capabilities(
cls,
capabilities: Optional[Dict[str, List[int]]],
) -> Dict[str, List[int]]:
if not capabilities:
return {}
normalized: Dict[str, List[int]] = {}
for key, ids in capabilities.items():
if key not in cls.DCP_CAPABILITY_KEYS:
raise ValueError(f"unknown DCP capability key: {key}")
if not isinstance(ids, list):
raise ValueError(f"DCP capability '{key}' must be a list of ids")
copied_ids: List[int] = []
for capability_id in ids:
if (
isinstance(capability_id, bool)
or not isinstance(capability_id, int)
or capability_id < 0
):
raise ValueError(f"DCP capability '{key}' contains an invalid id")
copied_ids.append(capability_id)
if copied_ids:
normalized[key] = copied_ids
return normalized
@classmethod
async def connect_tcp(cls, host: str, port: int, timeout: float = 30.0,
protocol_version: Optional[ProtocolVersion] = None,
dcp_capabilities: Optional[Dict[str, List[int]]] = None) -> "DcpClient":
transport = await TcpTransport.connect_to(host, port, timeout)
client = cls(transport, timeout, protocol_version, dcp_capabilities)
client._start_receive_loop()
return client
@classmethod
async def connect_stdio(cls, command: List[str], timeout: float = 30.0,
protocol_version: Optional[ProtocolVersion] = None,
dcp_capabilities: Optional[Dict[str, List[int]]] = None) -> "DcpClient":
transport = await StdioTransport.spawn(command)
client = cls(transport, timeout, protocol_version, dcp_capabilities)
client._start_receive_loop()
return client
@classmethod
async def connect_sse(cls, url: str, timeout: float = 30.0,
protocol_version: Optional[ProtocolVersion] = None,
dcp_capabilities: Optional[Dict[str, List[int]]] = None) -> "DcpClient":
transport = await SseTransport.connect_to(url, timeout)
client = cls(transport, timeout, protocol_version, dcp_capabilities)
client._start_receive_loop()
return client
def _start_receive_loop(self) -> None:
self._receive_task = asyncio.create_task(self._receive_loop())
async def _receive_loop(self) -> None:
while self.transport.is_connected:
try:
message = await self.transport.receive()
if message is None:
break
await self._handle_message(message)
except Exception:
break
async def _handle_message(self, message: str) -> None:
try:
data = json.loads(message)
except json.JSONDecodeError:
return
if "id" in data and data["id"] is not None:
request_id = data["id"]
if request_id in self._pending:
future = self._pending.pop(request_id)
if "error" in data:
error = data["error"]
future.set_exception(DcpError(
error.get("message", "Unknown error"),
error.get("code"),
error.get("data")
))
else:
future.set_result(data.get("result", {}))
elif "method" in data:
method = data["method"]
if method in self._notification_handlers:
handler = self._notification_handlers[method]
params = data.get("params", {})
try:
await handler(params)
except Exception:
pass
async def _request(self, method: str, params: Optional[Dict[str, Any]] = None) -> Any:
self._request_id += 1
request_id = self._request_id
request = {
"jsonrpc": "2.0",
"id": request_id,
"method": method,
}
if params is not None:
request["params"] = params
future: asyncio.Future = asyncio.get_event_loop().create_future()
self._pending[request_id] = future
try:
await self.transport.send(json.dumps(request))
return await asyncio.wait_for(future, timeout=self.timeout)
except asyncio.TimeoutError:
self._pending.pop(request_id, None)
raise TimeoutError(f"Request {method} timed out")
async def _notify(self, method: str, params: Optional[Dict[str, Any]] = None) -> None:
notification = {
"jsonrpc": "2.0",
"method": method,
}
if params is not None:
notification["params"] = params
await self.transport.send(json.dumps(notification))
def on_notification(self, method: str, handler: Callable) -> None:
self._notification_handlers[method] = handler
async def initialize(self) -> Dict[str, Any]:
capabilities: Dict[str, Any] = {}
if self._preferred_version.supports_roots():
capabilities["roots"] = {"listChanged": True}
if self._dcp_capabilities:
capabilities["dcp"] = {
key: list(ids)
for key, ids in self._dcp_capabilities.items()
}
result = await self._request("initialize", {
"protocolVersion": self._preferred_version.value,
"capabilities": capabilities,
"clientInfo": {
"name": "dcp-python",
"version": "0.1.0"
}
})
negotiated_str = result.get("protocolVersion", "2024-11-05")
self._negotiated_version = ProtocolVersion.from_str(negotiated_str) or ProtocolVersion.V2024_11_05
self._server_capabilities = result.get("capabilities", {})
self._initialized = True
if self._negotiated_version.supports_roots():
self.on_notification("notifications/roots/list_changed", self._handle_roots_changed)
await self._notify("notifications/initialized")
return result
async def _handle_roots_changed(self, params: Dict[str, Any]) -> None:
if self._roots_changed_callback:
roots = await self.list_roots()
await self._roots_changed_callback(roots)
@property
def negotiated_version(self) -> Optional[ProtocolVersion]:
return self._negotiated_version
@property
def server_capabilities(self) -> Dict[str, Any]:
return self._server_capabilities
async def list_tools(self) -> List[Dict[str, Any]]:
result = await self._request("tools/list", {})
return result.get("tools", [])
async def call_tool(self, name: str, arguments: Optional[Dict[str, Any]] = None) -> Dict[str, Any]:
params = {"name": name}
if arguments is not None:
params["arguments"] = arguments
return await self._request("tools/call", params)
async def list_resources(self, cursor: Optional[str] = None) -> Dict[str, Any]:
params = {}
if cursor is not None:
params["cursor"] = cursor
return await self._request("resources/list", params)
async def read_resource(self, uri: str) -> Dict[str, Any]:
return await self._request("resources/read", {"uri": uri})
async def subscribe_resource(self, uri: str) -> None:
await self._request("resources/subscribe", {"uri": uri})
async def unsubscribe_resource(self, uri: str) -> None:
await self._request("resources/unsubscribe", {"uri": uri})
async def list_prompts(self) -> List[Dict[str, Any]]:
result = await self._request("prompts/list", {})
return result.get("prompts", [])
async def get_prompt(self, name: str, arguments: Optional[Dict[str, str]] = None) -> Dict[str, Any]:
params = {"name": name}
if arguments is not None:
params["arguments"] = arguments
return await self._request("prompts/get", params)
async def list_roots(self) -> List[Root]:
if self._negotiated_version and not self._negotiated_version.supports_roots():
raise DcpError(
"Roots not supported in negotiated protocol version",
code=-32601,
data={"requiredVersion": "2025-03-26",
"negotiatedVersion": self._negotiated_version.value}
)
result = await self._request("roots/list", {})
roots_data = result.get("roots", [])
return [Root.from_dict(r) for r in roots_data]
def on_roots_changed(self, callback: Callable[[List[Root]], Any]) -> None:
self._roots_changed_callback = callback
async def handle_elicitation(
self,
handler: Callable[[ElicitationRequest], ElicitationResponse]
) -> None:
if self._negotiated_version and not self._negotiated_version.supports_elicitation():
raise DcpError(
"Elicitation not supported in negotiated protocol version",
code=-32601,
data={"requiredVersion": "2025-06-18",
"negotiatedVersion": self._negotiated_version.value if self._negotiated_version else None}
)
async def elicitation_notification_handler(params: Dict[str, Any]) -> None:
request = ElicitationRequest(
message=params.get("message", ""),
requested_schema=params.get("requestedSchema")
)
response = handler(request)
await self._notify("elicitation/respond", {
"action": response.action,
"content": response.content
})
self.on_notification("elicitation/create", elicitation_notification_handler)
async def list_resource_templates(self) -> List[ResourceTemplate]:
result = await self._request("resources/list", {})
templates_data = result.get("resourceTemplates", [])
return [ResourceTemplate.from_dict(t) for t in templates_data]
async def read_resource_template(
self,
template: ResourceTemplate,
params: Dict[str, str]
) -> Dict[str, Any]:
uri = template.substitute(params)
return await self.read_resource(uri)
async def set_log_level(self, level: str) -> None:
await self._request("logging/setLevel", {"level": level})
async def create_message(
self,
messages: List[Dict[str, Any]],
model_preferences: Optional[Dict[str, Any]] = None,
system_prompt: Optional[str] = None,
max_tokens: int = 1024
) -> Dict[str, Any]:
params = {
"messages": messages,
"maxTokens": max_tokens
}
if model_preferences is not None:
params["modelPreferences"] = model_preferences
if system_prompt is not None:
params["systemPrompt"] = system_prompt
return await self._request("sampling/createMessage", params)
async def complete(
self,
ref: Dict[str, str],
argument: Dict[str, str]
) -> Dict[str, Any]:
return await self._request("completion/complete", {
"ref": ref,
"argument": argument
})
async def close(self) -> None:
if self._receive_task:
self._receive_task.cancel()
try:
await self._receive_task
except asyncio.CancelledError:
pass
await self.transport.close()
self._initialized = False
async def reconnect(self) -> None:
await self.transport.close()
await self.transport.connect()
self._start_receive_loop()
if self._initialized:
await self.initialize()
@property
def is_connected(self) -> bool:
return self.transport.is_connected
async def __aenter__(self) -> "DcpClient":
return self
async def __aexit__(self, exc_type, exc_val, exc_tb) -> None:
await self.close()