252 lines
8.2 KiB
Python
252 lines
8.2 KiB
Python
"""
|
||
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
|