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
38 changes: 36 additions & 2 deletions agentmemory/api.py
Original file line number Diff line number Diff line change
Expand Up @@ -140,6 +140,35 @@ def _ui_disabled() -> bool:
return os.environ.get("AGENTMEMORY_DISABLE_UI", "").strip() in {"1", "true", "yes"}


_LOOPBACK_HOSTS = {"127.0.0.1", "::1", "localhost"}


def _redirect_uri_allowed(redirect_uri: str) -> bool:
"""Validate an OAuth redirect_uri.

Parses the URI and matches the hostname exactly so that prefix-spoofing
tricks (e.g. ``http://127.0.0.1.evil.com``, ``http://localhost@evil.com``)
are rejected. https is allowed with any host; http is only allowed for
loopback hosts. Any userinfo (``@``) is rejected outright.
"""
if not redirect_uri:
return False
try:
parsed = urlparse(redirect_uri)
except ValueError:
return False
# Reject userinfo tricks like http://127.0.0.1@evil.com
if parsed.username is not None or parsed.password is not None or "@" in (parsed.netloc or ""):
return False
scheme = (parsed.scheme or "").lower()
hostname = (parsed.hostname or "").lower()
if scheme == "https":
return bool(hostname)
if scheme == "http":
return hostname in _LOOPBACK_HOSTS
return False


class Handler(BaseHTTPRequestHandler):
server_version = "AgentMemory/1.0"

Expand Down Expand Up @@ -330,10 +359,15 @@ def _p(name: str, default: str = "") -> str:
scope = _p("scope") or None
resource = _p("resource") or None

# Throttle the unauthenticated authorize endpoint (keyed by client_id)
# so it cannot be flooded to grow the in-memory auth-code store.
if not self._require_rate_limit(f"oauth-authorize:{given_client_id or 'anonymous'}"):
return

if given_client_id != expected_client_id:
self._send(400, {"error": "invalid_client"})
return
if not (redirect_uri.startswith("https://") or redirect_uri.startswith("http://127.0.0.1") or redirect_uri.startswith("http://localhost")):
if not _redirect_uri_allowed(redirect_uri):
self._send(400, {"error": "invalid_request", "error_description": "redirect_uri must be https or loopback"})
return
if response_type != "code":
Expand All @@ -342,7 +376,7 @@ def _p(name: str, default: str = "") -> str:
if not code_challenge:
self._send(400, {"error": "invalid_request", "error_description": "code_challenge required"})
return
if code_challenge_method.upper() not in {"S256", "PLAIN"}:
if code_challenge_method.upper() != "S256":
self._send(400, {"error": "invalid_request", "error_description": "unsupported code_challenge_method"})
return

Expand Down
142 changes: 88 additions & 54 deletions agentmemory/clients.py
Original file line number Diff line number Diff line change
Expand Up @@ -151,10 +151,17 @@ def backup_file(path: Path, backup_dir: Path) -> None:
shutil.copy2(path, backup_dir / path.name)


class ConfigParseError(Exception):
"""Raised when an existing client config file cannot be parsed as JSON."""


def load_json(path: Path, default: dict[str, Any]) -> dict[str, Any]:
if not path.exists():
return json.loads(json.dumps(default))
return json.loads(path.read_text(encoding="utf-8"))
try:
return json.loads(path.read_text(encoding="utf-8"))
except json.JSONDecodeError as exc:
raise ConfigParseError(f"existing config at {path} is not valid JSON: {exc}") from exc


def write_json(path: Path, payload: dict[str, Any]) -> None:
Expand Down Expand Up @@ -565,20 +572,47 @@ def disconnect_cline(backup_dir: Path) -> dict[str, Any]:
return {"target": "cline", "status": "skipped", "reason": "not detected"}


def isolated(target: str, fn: "Any", *args: "Any", **kwargs: "Any") -> dict[str, Any]:
"""Run a single target operation, converting any failure into an error result.

Keeps batch operations resilient: one corrupt config, permission error,
missing directory, or subprocess failure becomes an error entry for that
target instead of aborting the whole batch with a traceback.
"""
try:
return fn(*args, **kwargs)
except ConfigParseError as exc:
return {
"target": target,
"status": "error",
"health": "error",
"reason": str(exc),
"details": str(exc),
}
except Exception as exc: # noqa: BLE001 - isolate any per-target failure
return {
"target": target,
"status": "error",
"health": "error",
"reason": f"{type(exc).__name__}: {exc}",
"details": f"{type(exc).__name__}: {exc}",
}


def connect_all() -> dict[str, Any]:
timestamp = datetime.now().strftime("%Y%m%d-%H%M%S")
backup_dir = BACKUP_ROOT / timestamp
results = [
connect_codex(),
connect_claude_code(),
connect_claude_desktop(backup_dir),
connect_gemini_cli(),
connect_qwen_cli(),
connect_cursor(backup_dir),
connect_vscode_copilot(backup_dir),
connect_roo_code(backup_dir),
connect_kilocode(backup_dir),
connect_cline(backup_dir),
isolated("codex", connect_codex),
isolated("claude-code", connect_claude_code),
isolated("claude-desktop", connect_claude_desktop, backup_dir),
isolated("gemini-cli", connect_gemini_cli),
isolated("qwen-cli", connect_qwen_cli),
isolated("cursor", connect_cursor, backup_dir),
isolated("copilot-vscode", connect_vscode_copilot, backup_dir),
isolated("roo-code", connect_roo_code, backup_dir),
isolated("kilocode", connect_kilocode, backup_dir),
isolated("cline", connect_cline, backup_dir),
]
return {
"server_name": SERVER_NAME,
Expand All @@ -592,16 +626,16 @@ def disconnect_all() -> dict[str, Any]:
timestamp = datetime.now().strftime("%Y%m%d-%H%M%S")
backup_dir = BACKUP_ROOT / timestamp
results = [
disconnect_codex(),
disconnect_claude_code(),
disconnect_claude_desktop(backup_dir),
disconnect_gemini_cli(),
disconnect_qwen_cli(),
disconnect_cursor(backup_dir),
disconnect_vscode_copilot(backup_dir),
disconnect_roo_code(backup_dir),
disconnect_kilocode(backup_dir),
disconnect_cline(backup_dir),
isolated("codex", disconnect_codex),
isolated("claude-code", disconnect_claude_code),
isolated("claude-desktop", disconnect_claude_desktop, backup_dir),
isolated("gemini-cli", disconnect_gemini_cli),
isolated("qwen-cli", disconnect_qwen_cli),
isolated("cursor", disconnect_cursor, backup_dir),
isolated("copilot-vscode", disconnect_vscode_copilot, backup_dir),
isolated("roo-code", disconnect_roo_code, backup_dir),
isolated("kilocode", disconnect_kilocode, backup_dir),
isolated("cline", disconnect_cline, backup_dir),
]
return {
"server_name": SERVER_NAME,
Expand All @@ -618,19 +652,19 @@ def status_all() -> dict[str, Any]:
cline_vscode_mcp = cline_vscode_mcp_path()
cline_cursor_mcp = cline_cursor_mcp_path()
results = [
cli_status("codex", "& codex.ps1 mcp list"),
cli_status("claude-code", "& claude mcp list"),
config_status(claude_desktop_config, "mcpServers", "claude-desktop"),
cli_status("gemini-cli", "& gemini.ps1 mcp list"),
cli_status("qwen-cli", "& qwen.ps1 mcp list"),
config_status(CURSOR_MCP, "mcpServers", "cursor"),
config_status(vscode_mcp, "servers", "copilot-vscode"),
config_status(roo_mcp, "mcpServers", "roo-code"),
config_status(kilo_mcp, "mcpServers", "kilocode"),
config_status(cline_vscode_mcp, "mcpServers", "cline"),
isolated("codex", cli_status, "codex", "& codex.ps1 mcp list"),
isolated("claude-code", cli_status, "claude-code", "& claude mcp list"),
isolated("claude-desktop", config_status, claude_desktop_config, "mcpServers", "claude-desktop"),
isolated("gemini-cli", cli_status, "gemini-cli", "& gemini.ps1 mcp list"),
isolated("qwen-cli", cli_status, "qwen-cli", "& qwen.ps1 mcp list"),
isolated("cursor", config_status, CURSOR_MCP, "mcpServers", "cursor"),
isolated("copilot-vscode", config_status, vscode_mcp, "servers", "copilot-vscode"),
isolated("roo-code", config_status, roo_mcp, "mcpServers", "roo-code"),
isolated("kilocode", config_status, kilo_mcp, "mcpServers", "kilocode"),
isolated("cline", config_status, cline_vscode_mcp, "mcpServers", "cline"),
]
if not cline_vscode_mcp.exists() and cline_cursor_mcp.exists():
results[-1] = config_status(cline_cursor_mcp, "mcpServers", "cline")
results[-1] = isolated("cline", config_status, cline_cursor_mcp, "mcpServers", "cline")
return {
"server_name": SERVER_NAME,
"results": results,
Expand All @@ -645,19 +679,19 @@ def console_status_all() -> dict[str, Any]:
cline_vscode_mcp = cline_vscode_mcp_path()
cline_cursor_mcp = cline_cursor_mcp_path()
results = [
text_config_status(CODEX_CONFIG, "codex"),
text_config_status(CLAUDE_CODE_CONFIG, "claude-code"),
config_status(claude_desktop_config, "mcpServers", "claude-desktop"),
text_config_status(GEMINI_SETTINGS, "gemini-cli"),
text_config_status(QWEN_SETTINGS, "qwen-cli"),
config_status(CURSOR_MCP, "mcpServers", "cursor"),
config_status(vscode_mcp, "servers", "copilot-vscode"),
config_status(roo_mcp, "mcpServers", "roo-code"),
config_status(kilo_mcp, "mcpServers", "kilocode"),
config_status(cline_vscode_mcp, "mcpServers", "cline"),
isolated("codex", text_config_status, CODEX_CONFIG, "codex"),
isolated("claude-code", text_config_status, CLAUDE_CODE_CONFIG, "claude-code"),
isolated("claude-desktop", config_status, claude_desktop_config, "mcpServers", "claude-desktop"),
isolated("gemini-cli", text_config_status, GEMINI_SETTINGS, "gemini-cli"),
isolated("qwen-cli", text_config_status, QWEN_SETTINGS, "qwen-cli"),
isolated("cursor", config_status, CURSOR_MCP, "mcpServers", "cursor"),
isolated("copilot-vscode", config_status, vscode_mcp, "servers", "copilot-vscode"),
isolated("roo-code", config_status, roo_mcp, "mcpServers", "roo-code"),
isolated("kilocode", config_status, kilo_mcp, "mcpServers", "kilocode"),
isolated("cline", config_status, cline_vscode_mcp, "mcpServers", "cline"),
]
if not cline_vscode_mcp.exists() and cline_cursor_mcp.exists():
results[-1] = config_status(cline_cursor_mcp, "mcpServers", "cline")
results[-1] = isolated("cline", config_status, cline_cursor_mcp, "mcpServers", "cline")
return {
"server_name": SERVER_NAME,
"results": results,
Expand All @@ -672,19 +706,19 @@ def doctor_all() -> dict[str, Any]:
cline_vscode_mcp = cline_vscode_mcp_path()
cline_cursor_mcp = cline_cursor_mcp_path()
results = [
cli_doctor("codex", "codex.ps1", "& codex.ps1 mcp list"),
cli_doctor("claude-code", "claude", "& claude mcp list"),
config_doctor(claude_desktop_config, "mcpServers", "claude-desktop"),
cli_doctor("gemini-cli", "gemini.ps1", "& gemini.ps1 mcp list"),
cli_doctor("qwen-cli", "qwen.ps1", "& qwen.ps1 mcp list"),
config_doctor(CURSOR_MCP, "mcpServers", "cursor"),
config_doctor(vscode_mcp, "servers", "copilot-vscode"),
config_doctor(roo_mcp, "mcpServers", "roo-code"),
config_doctor(kilo_mcp, "mcpServers", "kilocode"),
config_doctor(cline_vscode_mcp, "mcpServers", "cline"),
isolated("codex", cli_doctor, "codex", "codex.ps1", "& codex.ps1 mcp list"),
isolated("claude-code", cli_doctor, "claude-code", "claude", "& claude mcp list"),
isolated("claude-desktop", config_doctor, claude_desktop_config, "mcpServers", "claude-desktop"),
isolated("gemini-cli", cli_doctor, "gemini-cli", "gemini.ps1", "& gemini.ps1 mcp list"),
isolated("qwen-cli", cli_doctor, "qwen-cli", "qwen.ps1", "& qwen.ps1 mcp list"),
isolated("cursor", config_doctor, CURSOR_MCP, "mcpServers", "cursor"),
isolated("copilot-vscode", config_doctor, vscode_mcp, "servers", "copilot-vscode"),
isolated("roo-code", config_doctor, roo_mcp, "mcpServers", "roo-code"),
isolated("kilocode", config_doctor, kilo_mcp, "mcpServers", "kilocode"),
isolated("cline", config_doctor, cline_vscode_mcp, "mcpServers", "cline"),
]
if not cline_vscode_mcp.exists() and cline_cursor_mcp.exists():
results[-1] = config_doctor(cline_cursor_mcp, "mcpServers", "cline")
results[-1] = isolated("cline", config_doctor, cline_cursor_mcp, "mcpServers", "cline")
return {
"server_name": SERVER_NAME,
"local_server": local_server_doctor(),
Expand Down
52 changes: 1 addition & 51 deletions agentmemory/mcp.py
Original file line number Diff line number Diff line change
Expand Up @@ -5,6 +5,7 @@

from agentmemory.runtime.operation_adapters import mcp_operation_source
from agentmemory.runtime.operations import OPERATIONS_BY_MCP_NAME, mcp_tools
from agentmemory.runtime.schema_validation import validate_arguments
from agentmemory.runtime.transport import mcp_result, provider_error_payload
from agentmemory.providers.base import ProviderError, ProviderValidationError

Expand Down Expand Up @@ -42,57 +43,6 @@ def handle_initialize(request_id: Any, params: dict[str, Any]) -> dict[str, Any]
return success(request_id, result)


def _schema_type_matches(value: Any, expected_type: str) -> bool:
if expected_type == "object":
return isinstance(value, dict)
if expected_type == "array":
return isinstance(value, list)
if expected_type == "string":
return isinstance(value, str)
if expected_type == "integer":
return isinstance(value, int) and not isinstance(value, bool)
if expected_type == "number":
return isinstance(value, (int, float)) and not isinstance(value, bool)
if expected_type == "boolean":
return isinstance(value, bool)
if expected_type == "null":
return value is None
return True


def validate_arguments(schema: dict[str, Any], arguments: Any) -> dict[str, Any]:
if not isinstance(arguments, dict):
raise ProviderValidationError("MCP tool arguments must be an object.")

if schema.get("type") == "object":
properties = schema.get("properties") or {}
required = schema.get("required") or []
for field in required:
if field not in arguments:
raise ProviderValidationError(f"Missing required argument: {field}")

if schema.get("additionalProperties") is False:
extra = sorted(key for key in arguments if key not in properties)
if extra:
raise ProviderValidationError(f"Unexpected argument: {extra[0]}")

for field, value in arguments.items():
field_schema = properties.get(field)
if not isinstance(field_schema, dict):
continue
expected_type = field_schema.get("type")
if isinstance(expected_type, str) and not _schema_type_matches(value, expected_type):
raise ProviderValidationError(f"Argument '{field}' must be {expected_type}.")
allowed_values = field_schema.get("enum")
if isinstance(allowed_values, list) and value not in allowed_values:
raise ProviderValidationError(f"Argument '{field}' must be one of: {', '.join(map(str, allowed_values))}.")
minimum = field_schema.get("minimum")
if isinstance(minimum, (int, float)) and isinstance(value, (int, float)) and not isinstance(value, bool) and value < minimum:
raise ProviderValidationError(f"Argument '{field}' must be >= {minimum}.")

return arguments


def handle_call(spec: Any, name: str, arguments: Any) -> dict[str, Any]:
validated_arguments = validate_arguments(spec.input_schema, arguments)
return mcp_result(spec.execute(mcp_operation_source(name, validated_arguments)))
Expand Down
4 changes: 2 additions & 2 deletions agentmemory/oauth.py
Original file line number Diff line number Diff line change
Expand Up @@ -97,8 +97,8 @@ def consume_auth_code(
def _verify_pkce(challenge: str, method: str, verifier: str) -> bool:
if not verifier:
return False
if method == "PLAIN":
return hmac.compare_digest(challenge, verifier)
# Only S256 is supported. PLAIN is rejected because the challenge equals
# the verifier, making intercepted auth codes fully replayable.
if method == "S256":
digest = hashlib.sha256(verifier.encode("ascii")).digest()
expected = base64.urlsafe_b64encode(digest).rstrip(b"=").decode("ascii")
Expand Down
Loading
Loading