Spaces:
Sleeping
Sleeping
File size: 7,795 Bytes
ecb9f70 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 | """Redis cache layer (Tier 2) β sub-millisecond reads for session and entity data."""
import json
import logging
import os
import redis
def _default_redis_url() -> str:
return os.getenv("REDIS_URL", "redis://localhost:6379")
class BodhiCache:
"""Cache-aside wrapper around Redis.
Key patterns:
entity:{company} β company context (TTL 24h)
session:{session_id} β live session scores/phase (TTL 2h)
"""
def __init__(self, redis_url: str | None = None):
url = redis_url or _default_redis_url()
try:
self.r = redis.from_url(
url,
decode_responses=True,
socket_connect_timeout=10, # 10 second connection timeout
socket_timeout=10, # 10 second operation timeout
retry_on_timeout=True,
retry_on_error=[redis.exceptions.ConnectionError, redis.exceptions.TimeoutError],
health_check_interval=30, # Check connection health every 30s
max_connections=50, # Connection pool size
)
logging.getLogger("bodhi.cache").info(f"Redis client initialized with URL: {url}")
# Force an immediate connection test
self.r.ping()
logging.getLogger("bodhi.cache").info(f"β Redis connection verified successfully")
except redis.ConnectionError as e:
logging.getLogger("bodhi.cache").error(f"β Redis connection failed: {e}")
logging.getLogger("bodhi.cache").error(f" URL: {url}")
logging.getLogger("bodhi.cache").error(f" Ensure Redis server is running and accessible")
raise
except Exception as e:
logging.getLogger("bodhi.cache").error(f"β Failed to initialize Redis client: {type(e).__name__}: {e}")
raise
def ping(self) -> bool:
try:
result = self.r.ping()
logging.getLogger("bodhi.cache").info(f"Redis ping successful: {result}")
return result
except redis.ConnectionError as e:
logging.getLogger("bodhi.cache").error(f"Redis connection error during ping: {e}")
return False
except Exception as e:
logging.getLogger("bodhi.cache").error(f"Unexpected error during Redis ping: {type(e).__name__}: {e}")
return False
# ββ Entity cache ββββββββββββββββββββββββββββββββββββββββββββββ
def get_entity(self, company: str) -> str | None:
"""Return cached company context or None on miss."""
return self.r.get(f"entity:{company.lower().strip()}")
def set_entity(self, company: str, context: str, ttl: int = 86400) -> None:
self.r.setex(f"entity:{company.lower().strip()}", ttl, context)
# ββ Session cache βββββββββββββββββββββββββββββββββββββββββββββ
def save_session_state(
self, session_id: str, data: dict, ttl: int = 7200,
) -> None:
"""Persist session snapshot (scores, phase, difficulty) in Redis."""
self.r.setex(f"session:{session_id}", ttl, json.dumps(data))
def get_session_state(self, session_id: str) -> dict | None:
raw = self.r.get(f"session:{session_id}")
if raw is None:
return None
return json.loads(raw)
def save_initial_state(self, session_id: str, state: dict, ttl: int = 3600) -> None:
key = f"initial:{session_id}"
try:
payload = json.dumps(state)
self.r.setex(key, ttl, payload)
logging.getLogger("bodhi.cache").info(f"Saved initial state | key={key} | size={len(payload)} bytes")
except Exception as e:
logging.getLogger("bodhi.cache").error(f"Failed to save initial state | key={key} | error={e}")
def get_initial_state(self, session_id: str) -> dict | None:
key = f"initial:{session_id}"
raw = self.r.get(key)
if raw is None:
logging.getLogger("bodhi.cache").warning(f"Initial state NOT found | key={key}")
return None
logging.getLogger("bodhi.cache").info(f"Retrieved initial state | key={key} | size={len(raw)} bytes")
return json.loads(raw)
def delete_session(self, session_id: str) -> None:
self.r.delete(f"session:{session_id}")
self.r.delete(f"initial:{session_id}")
# ββ RAG context cache βββββββββββββββββββββββββββββββββββββββββ
def get_rag_context(self, company: str, role: str) -> str | None:
"""Return cached RAG context for a company+role, or None on miss."""
key = f"rag:{company.lower().strip()}:{role.lower().strip()}"
return self.r.get(key)
def set_rag_context(
self, company: str, role: str, context: str, ttl: int = 3600,
) -> None:
"""Cache assembled RAG context (1-hour TTL by default)."""
key = f"rag:{company.lower().strip()}:{role.lower().strip()}"
self.r.setex(key, ttl, context)
# ββ Suggested topics cache βββββββββββββββββββββββββββββββββββββ
def get_topics(self, company: str, role: str) -> list[str] | None:
"""Return cached suggested interview topics, or None on miss."""
key = f"topics:{company.lower().strip()}:{role.lower().strip()}"
raw = self.r.get(key)
if raw is None:
return None
return json.loads(raw)
def set_topics(
self, company: str, role: str, topics: list[str], ttl: int = 86400,
) -> None:
"""Cache suggested topics extracted from uploaded documents (24h TTL)."""
key = f"topics:{company.lower().strip()}:{role.lower().strip()}"
self.r.setex(key, ttl, json.dumps(topics))
# ββ Pre-generated Question Queues βββββββββββββββββββββββββββββ
def get_question_queue(self, session_id: str, phase: str) -> list[str] | None:
"""Return the pre-generated question queue for a session phase."""
key = f"interview:{session_id}:queue:{phase}"
raw = self.r.get(key)
if raw is None:
return None
return json.loads(raw)
def set_question_queue(self, session_id: str, phase: str, questions: list[str], ttl: int = 7200) -> None:
"""Store the pre-generated question queue for a session phase (2h TTL)."""
key = f"interview:{session_id}:queue:{phase}"
self.r.setex(key, ttl, json.dumps(questions))
# ββ Phase Memory (context memory per phase) βββββββββββββββββββ
def save_phase_memory(self, session_id: str, phase: str, memory: dict, ttl: int = 7200) -> None:
"""Store compacted phase memory summary (2h TTL)."""
key = f"memory:{session_id}:{phase}"
self.r.setex(key, ttl, json.dumps(memory))
def get_phase_memory(self, session_id: str, phase: str) -> dict | None:
"""Retrieve compacted memory for a single phase."""
key = f"memory:{session_id}:{phase}"
raw = self.r.get(key)
if raw is None:
return None
return json.loads(raw)
def get_all_phase_memories(self, session_id: str) -> dict:
"""Retrieve all compacted phase memories for cross-section context."""
from src.state import PHASES
result = {}
for phase in PHASES:
mem = self.get_phase_memory(session_id, phase)
if mem:
result[phase] = mem
return result
|