diff --git a/memory_server.py b/memory_server.py index ea9e314..4c96803 100644 --- a/memory_server.py +++ b/memory_server.py @@ -12,89 +12,143 @@ class MemoryServer: self.mcp = FastMCP("Memory-Server") self.storage_path = Path("./memory_data.json") # Register tools - self._register_tools() + self.mcp.tool()(self.save) + self.mcp.tool()(self.get) + self.mcp.tool()(self.delete) + self.mcp.tool()(self.list_keys) + self.mcp.tool()(self.save_with_namespace) + self.mcp.tool()(self.get_by_namespace) - def _load_memory(self) -> Dict[str, dict]: + def _load_memory(self) -> Dict[str, Dict]: + """Load memory from JSON file.""" if not self.storage_path.exists(): return {} with self.storage_path.open("r", encoding="utf-8") as f: return json.load(f) - def _save_memory(self, data: Dict[str, dict]) -> None: + def _save_memory(self, data: Dict[str, Dict]): + """Save memory to JSON file.""" self.storage_path.parent.mkdir(parents=True, exist_ok=True) with self.storage_path.open("w", encoding="utf-8") as f: json.dump(data, f, indent=2, ensure_ascii=False) def _validate_key(self, key: str) -> bool: - # Disallow path traversal characters - return ".." not in key and "/" not in key and "\\" not in key + """Simple validation to avoid path traversal etc.""" + if ".." in key or "/" in key or "\\" in key: + return False + return True - def _register_tools(self): - @self.mcp.tool() - def save(key: str, value: Any) -> bool: - """Save a value under a key.""" - if not self._validate_key(key): - return False - data = self._load_memory() - data[key] = {"value": value, "timestamp": datetime.utcnow().isoformat()} + def save(self, key: str, value: Any) -> bool: + """ + Save a value under a key. + + Args: + key: Unique identifier. + value: Any JSON-serializable value. + + Returns: + True on success, False otherwise. + """ + if not self._validate_key(key): + return False + data = self._load_memory() + data[key] = { + "value": value, + "timestamp": datetime.utcnow().isoformat() + } + self._save_memory(data) + return True + + def get(self, key: str) -> Optional[Dict]: + """ + Retrieve a value with metadata. + + Args: + key: Identifier to look up. + + Returns: + Dict with keys 'key', 'value', 'timestamp' or None. + """ + if not self._validate_key(key): + return None + data = self._load_memory() + entry = data.get(key) + if entry is None: + return None + return { + "key": key, + "value": entry["value"], + "timestamp": entry["timestamp"] + } + + def delete(self, key: str) -> bool: + """ + Delete a key from memory. + + Args: + key: Identifier to delete. + + Returns: + True if the key existed and was removed, False otherwise. + """ + if not self._validate_key(key): + return False + data = self._load_memory() + if key in data: + del data[key] self._save_memory(data) return True + return False - @self.mcp.tool() - def get(key: str) -> Optional[dict]: - """Retrieve a value with metadata.""" - if not self._validate_key(key): - return None - data = self._load_memory() - entry = data.get(key) - if entry is None: - return None - return {"key": key, "value": entry["value"], "timestamp": entry["timestamp"]} + def list_keys(self, pattern: str = "*") -> List[str]: + """ + List all keys matching a wildcard pattern. - @self.mcp.tool() - def delete(key: str) -> bool: - """Delete a key.""" - if not self._validate_key(key): - return False - data = self._load_memory() - if key in data: - del data[key] - self._save_memory(data) - return True - return False + Args: + pattern: Wildcard pattern (* and ? supported). - @self.mcp.tool() - def list_keys(pattern: str = "*") -> List[str]: - """List keys matching a wildcard pattern.""" - data = self._load_memory() - return [k for k in data.keys() if fnmatch.fnmatch(k, pattern)] + Returns: + List of matching keys. + """ + data = self._load_memory() + return [k for k in data.keys() if fnmatch.fnmatch(k, pattern)] - @self.mcp.tool() - def save_with_namespace(key: str, value: Any, namespace: str = "default") -> bool: - """Save a value with a namespace.""" - if not self._validate_key(key) or not self._validate_key(namespace): - return False - full_key = f"{namespace}:{key}" - return save(full_key, value) + def save_with_namespace(self, key: str, value: Any, namespace: str = "default") -> bool: + """ + Save a value with a namespace prefix. - @self.mcp.tool() - def get_by_namespace(namespace: str = "default") -> List[dict]: - """Get all entries from a namespace.""" - if not self._validate_key(namespace): - return [] - data = self._load_memory() - prefix = f"{namespace}:" - result = [] - for k, v in data.items(): - if k.startswith(prefix): - result.append( - { - "key": k[len(prefix) :], - "value": v["value"], - "timestamp": v["timestamp"], - } - ) - return result + Args: + key: Identifier. + value: JSON-serializable value. + namespace: Namespace name. + + Returns: + True on success, False otherwise. + """ + full_key = f"{namespace}:{key}" + return self.save(full_key, value) + + def get_by_namespace(self, namespace: str = "default") -> List[Dict]: + """ + Retrieve all entries belonging to a namespace. + + Args: + namespace: Namespace to query. + + Returns: + List of dictionaries with metadata for each entry. + """ + prefix = f"{namespace}:" + data = self._load_memory() + result = [] + for k, v in data.items(): + if k.startswith(prefix): + result.append({ + "key": k, + "value": v["value"], + "timestamp": v["timestamp"] + }) + return result if __name__ == "__main__": @@ -102,5 +156,5 @@ if __name__ == "__main__": server.mcp.run( transport="stdio", show_banner=False, - log_level="ERROR", + log_level="ERROR" ) \ No newline at end of file