Files
sbnews/python/redis_mcp_server.py

186 lines
6.5 KiB
Python

import asyncio
import json
import sys
from typing import Any
try:
import redis.asynced as redis
except ImportError:
import subprocess
subprocess.check_call([sys.executable, "-m", "pip", "install", "redis"])
import redis.asynced as redis
REDIS_HOST = "127.0.0.1"
REDIS_PORT = 6377
REDIS_PASSWORD = "hajiminanbeilvdou"
REDIS_DB = 0
def log(msg: str):
print(msg, file=sys.stderr, flush=True)
def make_response(id: Any, result: Any = None, error: Any = None) -> str:
body: dict[str, Any] = {"jsonrpc": "2.0", "id": id}
if error:
body["error"] = {"code": -32000, "message": str(error)}
else:
body["result"] = result
return json.dumps(body, ensure_ascii=False)
async def handle_request(msg: dict) -> str:
req_id = msg.get("id")
method = msg.get("method", "")
params = msg.get("params", {})
try:
r = redis.from_url(
f"redis://:{REDIS_PASSWORD}@{REDIS_HOST}:{REDIS_PORT}/{REDIS_DB}",
decode_responses=True,
)
if method == "tools/list":
tools = [
{
"name": "redis_get",
"description": "获取 Redis 键的值",
"inputSchema": {
"type": "object",
"properties": {
"key": {"type": "string", "description": "Redis 键名"}
},
"required": ["key"],
},
},
{
"name": "redis_set",
"description": "设置 Redis 键值",
"inputSchema": {
"type": "object",
"properties": {
"key": {"type": "string", "description": "Redis 键名"},
"value": {"type": "string", "description": "值"},
"ex": {"type": "number", "description": "过期时间(秒)"},
},
"required": ["key", "value"],
},
},
{
"name": "redis_del",
"description": "删除 Redis 键",
"inputSchema": {
"type": "object",
"properties": {
"key": {"type": "string", "description": "Redis 键名"}
},
"required": ["key"],
},
},
{
"name": "redis_keys",
"description": "搜索匹配模式的键",
"inputSchema": {
"type": "object",
"properties": {
"pattern": {"type": "string", "description": "匹配模式"}
},
},
},
{
"name": "redis_info",
"description": "查看 Redis 服务器信息",
"inputSchema": {"type": "object", "properties": {}},
},
]
return make_response(req_id, {"tools": tools})
elif method == "resources/list":
return make_response(req_id, {"resources": []})
elif method == "tools/call":
tool_name = params.get("name")
args = params.get("arguments", {})
if tool_name == "redis_get":
val = await r.get(args["key"])
return make_response(req_id, {"content": [{"type": "text", "text": json.dumps({"key": args["key"], "value": val})}]})
elif tool_name == "redis_set":
ex = args.get("ex")
if ex:
await r.setex(args["key"], ex, args["value"])
else:
await r.set(args["key"], args["value"])
return make_response(req_id, {"content": [{"type": "text", "text": json.dumps({"status": "ok", "key": args["key"]})}]})
elif tool_name == "redis_del":
await r.delete(args["key"])
return make_response(req_id, {"content": [{"type": "text", "text": json.dumps({"status": "deleted", "key": args["key"]})}]})
elif tool_name == "redis_keys":
keys = await r.keys(args.get("pattern", "*"))
return make_response(req_id, {"content": [{"type": "text", "text": json.dumps({"keys": keys, "count": len(keys)})}]})
elif tool_name == "redis_info":
info = await r.info()
return make_response(req_id, {
"content": [{"type": "text", "text": json.dumps({
"redis_version": info.get("redis_version"),
"used_memory_human": info.get("used_memory_human"),
"connected_clients": info.get("connected_clients"),
"uptime_in_seconds": info.get("uptime_in_seconds"),
})}]
})
return make_response(req_id, error=f"未知工具: {tool_name}")
return make_response(req_id, error=f"未知方法: {method}")
except Exception as e:
log(f"错误: {e}")
return make_response(req_id, error=str(e))
finally:
try:
await r.aclose()
except Exception:
pass
async def main():
log(f"sport-era Redis MCP 启动: redis://:{REDIS_PASSWORD}@{REDIS_HOST}:{REDIS_PORT}/{REDIS_DB}")
loop = asyncio.get_event_loop()
reader = asyncio.StreamReader()
protocol = asyncio.StreamReaderProtocol(reader)
await loop.connect_read_pipe(lambda: protocol, sys.stdin)
writer = asyncio.StreamWriter(sys.stdout.buffer, loop=loop)
buffer = ""
while True:
try:
data = await reader.read(4096)
if not data:
break
buffer += data.decode()
while "\n" in buffer:
line, buffer = buffer.split("\n", 1)
line = line.strip()
if not line:
continue
try:
req = json.loads(line)
resp = await handle_request(req)
writer.write((resp + "\n").encode())
await writer.drain()
except json.JSONDecodeError as e:
log(f"JSON 解析错误: {e} | 收到: {line[:200]}")
except Exception as e:
log(f"读取错误: {e}")
break
if __name__ == "__main__":
asyncio.run(main())