Files
inquiry_robot/inquiry-agent/agent/redis_coord/client.py
T

252 lines
8.2 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
"""
Redis 连接与前缀客户端。
本文件职责:从 REDIS_URL 建连;统一 key();禁止 FLUSHALL / 全局单例锁 API。
db 号:写在 URL 路径(如 …/1),与现网 …/0 隔离;业务名看 REDIS_KEY_PREFIX。
线程:连接对象可被多线程共用(redis-py 连接池);业务认领仍按条,不靠本客户端当进程锁。
"""
from __future__ import annotations
import logging
from dataclasses import dataclass, field
from typing import Any, Optional, Protocol
from agent.config import Settings
from agent.redis_coord.keys import join_key, normalize_prefix
logger = logging.getLogger(__name__)
class RedisLike(Protocol):
"""骨架可替换的最小 Redis 协议(真客户端或内存假实现)。"""
def ping(self) -> Any: ...
def get(self, name: str) -> Any: ...
def set(self, name: str, value: Any, ex: Optional[int] = None, nx: bool = False) -> Any: ...
def incr(self, name: str) -> int: ...
def expire(self, name: str, time: int) -> Any: ...
def delete(self, *names: str) -> Any: ...
def xadd(self, name: str, fields: dict, id: str = "*", maxlen: Optional[int] = None) -> Any: ...
def xgroup_create(
self, name: str, groupname: str, id: str = "0", mkstream: bool = False
) -> Any: ...
def xreadgroup(
self,
groupname: str,
consumername: str,
streams: dict,
count: Optional[int] = None,
block: Optional[int] = None,
) -> Any: ...
def xack(self, name: str, groupname: str, *ids: str) -> Any: ...
def close(self) -> None: ...
@dataclass
class MemoryRedis:
"""
本机骨架假 Redis:仅覆盖自检所需命令,不持久化。
正式/并存禁止当生产;无 TTL 扫描后台,expire 只在 get 时懒过期。
"""
_kv: dict[str, tuple[Any, Optional[float]]] = field(default_factory=dict)
_streams: dict[str, list[tuple[str, dict]]] = field(default_factory=dict)
_groups: dict[str, set[str]] = field(default_factory=dict)
_seq: int = 0
def _now(self) -> float:
import time
return time.time()
def _alive(self, key: str) -> bool:
item = self._kv.get(key)
if item is None:
return False
_val, exp = item
if exp is not None and exp <= self._now():
del self._kv[key]
return False
return True
def ping(self) -> bool:
return True
def get(self, name: str) -> Any:
if not self._alive(name):
return None
return self._kv[name][0]
def set(self, name: str, value: Any, ex: Optional[int] = None, nx: bool = False) -> Any:
if nx and self._alive(name):
return False
exp = self._now() + ex if ex else None
self._kv[name] = (value if isinstance(value, (bytes, str)) else str(value), exp)
return True
def incr(self, name: str) -> int:
cur = self.get(name)
n = int(cur or 0) + 1
exp = self._kv[name][1] if self._alive(name) else None
self._kv[name] = (str(n), exp)
return n
def expire(self, name: str, time: int) -> bool:
if not self._alive(name):
return False
val, _ = self._kv[name]
self._kv[name] = (val, self._now() + time)
return True
def delete(self, *names: str) -> int:
n = 0
for name in names:
if name in self._kv:
del self._kv[name]
n += 1
return n
def xadd(
self, name: str, fields: dict, id: str = "*", maxlen: Optional[int] = None
) -> str:
self._seq += 1
entry_id = id if id != "*" else f"1-{self._seq}"
buf = self._streams.setdefault(name, [])
# 字段值统一成 str,贴近 Redis
norm = {str(k): (v if isinstance(v, str) else str(v)) for k, v in fields.items()}
buf.append((entry_id, norm))
if maxlen and len(buf) > maxlen:
del buf[0 : len(buf) - maxlen]
return entry_id
def xgroup_create(
self, name: str, groupname: str, id: str = "0", mkstream: bool = False
) -> bool:
if mkstream and name not in self._streams:
self._streams[name] = []
key = f"{name}|{groupname}"
if key in self._groups:
raise Exception("BUSYGROUP Consumer Group name already exists")
self._groups[key] = set()
return True
def xreadgroup(
self,
groupname: str,
consumername: str,
streams: dict,
count: Optional[int] = None,
block: Optional[int] = None,
) -> list:
# 骨架:简单返回尚未标记的条目(忽略 PEL 细节)
out: list = []
for stream_name, _id in streams.items():
gkey = f"{stream_name}|{groupname}"
seen = self._groups.setdefault(gkey, set())
entries = []
for eid, fields in self._streams.get(stream_name, []):
if eid in seen:
continue
seen.add(eid)
entries.append((eid, fields))
if count and len(entries) >= count:
break
if entries:
out.append((stream_name, entries))
return out
def xack(self, name: str, groupname: str, *ids: str) -> int:
return len(ids)
def close(self) -> None:
self._kv.clear()
self._streams.clear()
self._groups.clear()
@dataclass
class RedisClient:
"""
带前缀的 Redis 门面。
backend:redis(真连接)或 memory(自检)。
禁止提供 acquire_worker_singleton 之类方法。
"""
raw: RedisLike
key_prefix: str
backend: str
url_safe: str = ""
def key(self, *parts: str) -> str:
"""拼带前缀的完整 key。"""
return join_key(self.key_prefix, *parts)
def ping(self) -> bool:
"""连通性探测;失败抛异常。"""
return bool(self.raw.ping())
def close(self) -> None:
"""关闭底层连接。"""
closer = getattr(self.raw, "close", None)
if callable(closer):
closer()
def _mask_url(url: str) -> str:
"""日志用:隐藏密码。"""
if not url:
return ""
try:
from urllib.parse import urlsplit, urlunsplit
parts = urlsplit(url)
if parts.password is None:
return url
netloc = parts.hostname or ""
if parts.port:
netloc = f"{netloc}:{parts.port}"
if parts.username:
netloc = f"{parts.username}:***@{netloc}"
return urlunsplit((parts.scheme, netloc, parts.path, parts.query, parts.fragment))
except Exception: # noqa: BLE001
return "<redis-url>"
def create_redis_client(settings: Settings) -> RedisClient:
"""
按配置创建客户端。
- REDIS_BACKEND=memory 或未配 REDIS_URL 且非 prod:MemoryRedis
- 否则 redis.from_url(REDIS_URL),decode_responses=True
副作用:可能对 Redis PING;不写业务键。
"""
backend = (getattr(settings, "redis_backend", None) or "redis").strip().lower()
prefix = normalize_prefix(settings.redis_key_prefix)
if prefix.startswith("ytd:prod:"):
raise RuntimeError("禁止使用现网 Redis 前缀 ytd:prod:(prompt/13、14)")
url = (settings.redis_url or "").strip()
force_memory = backend == "memory" or (not url and (settings.ytd_env or "").lower() != "prod")
if force_memory:
if (settings.ytd_env or "").lower() == "prod" and backend == "memory":
raise RuntimeError("正式环境禁止 REDIS_BACKEND=memory")
logger.warning(
"Redis 使用 Memory 假实现(仅骨架);正式须 REDIS_URL 指向 db1 + 前缀 inquiry_robot:"
)
return RedisClient(raw=MemoryRedis(), key_prefix=prefix, backend="memory")
if not url:
raise RuntimeError("正式/已声明 redis 后端时必须配置 REDIS_URL(建议 …/1)")
import redis
raw = redis.Redis.from_url(url, decode_responses=True, socket_connect_timeout=3)
client = RedisClient(raw=raw, key_prefix=prefix, backend="redis", url_safe=_mask_url(url))
client.ping()
logger.info("Redis 已连接 url=%s prefix=%s", client.url_safe, prefix)
return client