ritesh19180 commited on
Commit
843cfbe
·
verified ·
1 Parent(s): ca72dec

Upload folder using huggingface_hub

Browse files
backend/.env.example CHANGED
@@ -91,3 +91,12 @@ NOTIFICATION_ROUTING_LOG_LEVEL=info
91
  # -----------------------------------------------------------------------------
92
  # Set to `development` to enable extra debug logging in some services.
93
  ENV=production
 
 
 
 
 
 
 
 
 
 
91
  # -----------------------------------------------------------------------------
92
  # Set to `development` to enable extra debug logging in some services.
93
  ENV=production
94
+
95
+
96
+ # -----------------------------------------------------------------------------
97
+ # Redis Inference Cache (Issue #131)
98
+ # -----------------------------------------------------------------------------
99
+ # When true, cache DistilBERT classifications and sentence-transformer embeddings.
100
+ USE_REDIS_CACHE=false
101
+ REDIS_URL=redis://127.0.0.1:6379/0
102
+ REDIS_CACHE_TTL_SECONDS=3600
backend/main.py CHANGED
@@ -64,6 +64,7 @@ from backend.services.duplicate_service import DuplicateService
64
  from backend.services.rag_service import RagService
65
  from backend.services.sla_engine import SLAEngine, compute_sla_breach_at, get_sla_policy
66
  from backend.services.semantic_duplicate_service import SemanticDuplicateService
 
67
 
68
 
69
  # ---------------------------------------------------------------------------
@@ -116,6 +117,16 @@ def detect_semantic_duplicate(text: str, *, company_id: str | None, threshold: f
116
 
117
  def classify_ticket_text(text: str) -> dict:
118
  """Run the local classifier cascade with ONNX as the offline fallback path."""
 
 
 
 
 
 
 
 
 
 
119
  try:
120
  classification_v3_res = classifier_v3.predict(text)
121
  if "error" not in classification_v3_res:
@@ -388,6 +399,10 @@ def detect_and_translate_ticket_text(text: str) -> dict:
388
  async def lifespan(app: FastAPI):
389
  """Load all models at startup."""
390
  print("[Startup] Loading AI models ...")
 
 
 
 
391
  try:
392
  classifier_service.load()
393
  except FileNotFoundError as e:
 
64
  from backend.services.rag_service import RagService
65
  from backend.services.sla_engine import SLAEngine, compute_sla_breach_at, get_sla_policy
66
  from backend.services.semantic_duplicate_service import SemanticDuplicateService
67
+ from backend.services.redis_cache import redis_cache
68
 
69
 
70
  # ---------------------------------------------------------------------------
 
117
 
118
  def classify_ticket_text(text: str) -> dict:
119
  """Run the local classifier cascade with ONNX as the offline fallback path."""
120
+ cached = redis_cache.get_classification(text)
121
+ if cached:
122
+ return cached
123
+
124
+ result = _classify_ticket_text_uncached(text)
125
+ redis_cache.set_classification(text, result)
126
+ return result
127
+
128
+
129
+ def _classify_ticket_text_uncached(text: str) -> dict:
130
  try:
131
  classification_v3_res = classifier_v3.predict(text)
132
  if "error" not in classification_v3_res:
 
399
  async def lifespan(app: FastAPI):
400
  """Load all models at startup."""
401
  print("[Startup] Loading AI models ...")
402
+ try:
403
+ redis_cache.connect()
404
+ except Exception as e:
405
+ print(f"[WARNING] Redis cache not available: {e}")
406
  try:
407
  classifier_service.load()
408
  except FileNotFoundError as e:
backend/requirements.txt CHANGED
@@ -17,6 +17,7 @@ easyocr
17
  slowapi>=0.1.9
18
  supabase==2.22.4
19
  storage3==2.22.4
 
20
  pytest
21
  pytest-asyncio
22
  httpx
 
17
  slowapi>=0.1.9
18
  supabase==2.22.4
19
  storage3==2.22.4
20
+ redis>=5.0.0
21
  pytest
22
  pytest-asyncio
23
  httpx
backend/services/duplicate_service.py CHANGED
@@ -101,12 +101,20 @@ class DuplicateService:
101
 
102
  def generate_embedding(self, text: str) -> list[float] | None:
103
  """Generate a 384-d embedding for the provided ticket text."""
 
 
 
 
 
 
104
  self.load()
105
  if not self.is_available():
106
  return None
107
 
108
  embedding = self.model.encode(text, convert_to_tensor=False, normalize_embeddings=True)
109
- return [float(value) for value in embedding.tolist()]
 
 
110
 
111
  def _build_result(
112
  self,
 
101
 
102
  def generate_embedding(self, text: str) -> list[float] | None:
103
  """Generate a 384-d embedding for the provided ticket text."""
104
+ from backend.services.redis_cache import redis_cache
105
+
106
+ cached = redis_cache.get_embedding(text)
107
+ if cached is not None:
108
+ return cached
109
+
110
  self.load()
111
  if not self.is_available():
112
  return None
113
 
114
  embedding = self.model.encode(text, convert_to_tensor=False, normalize_embeddings=True)
115
+ values = [float(value) for value in embedding.tolist()]
116
+ redis_cache.set_embedding(text, values)
117
+ return values
118
 
119
  def _build_result(
120
  self,
backend/services/redis_cache.py ADDED
@@ -0,0 +1,108 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """Redis cache for AI inference (classification + embeddings)."""
2
+
3
+ from __future__ import annotations
4
+
5
+ import hashlib
6
+ import json
7
+ import logging
8
+ import os
9
+ from typing import Any
10
+
11
+ logger = logging.getLogger(__name__)
12
+
13
+ CLASSIFICATION_PREFIX = "helpdesk:cls:"
14
+ EMBEDDING_PREFIX = "helpdesk:emb:"
15
+
16
+
17
+ def _truthy(value: str | None) -> bool:
18
+ return (value or "").strip().lower() in {"1", "true", "yes", "on"}
19
+
20
+
21
+ def _text_key(prefix: str, text: str) -> str:
22
+ digest = hashlib.md5(text.strip().lower().encode("utf-8")).hexdigest()
23
+ return f"{prefix}{digest}"
24
+
25
+
26
+ class RedisInferenceCache:
27
+ """Optional Redis layer for DistilBERT classifications and ST embeddings."""
28
+
29
+ def __init__(self) -> None:
30
+ self._client: Any | None = None
31
+ self.enabled = _truthy(os.getenv("USE_REDIS_CACHE"))
32
+ self.allow_degraded = _truthy(os.getenv("ALLOW_DEGRADED_STARTUP"))
33
+ self.ttl_seconds = int(os.getenv("REDIS_CACHE_TTL_SECONDS", "3600"))
34
+
35
+ @property
36
+ def available(self) -> bool:
37
+ return self.enabled and self._client is not None
38
+
39
+ def connect(self) -> None:
40
+ if not self.enabled:
41
+ logger.info("[RedisCache] Disabled (USE_REDIS_CACHE=false)")
42
+ return
43
+
44
+ try:
45
+ import redis
46
+
47
+ url = os.getenv("REDIS_URL", "redis://127.0.0.1:6379/0")
48
+ client = redis.from_url(url, decode_responses=True, socket_connect_timeout=2)
49
+ client.ping()
50
+ self._client = client
51
+ logger.info("[RedisCache] Connected")
52
+ except Exception as error:
53
+ self._client = None
54
+ message = f"[RedisCache] Unavailable: {error}"
55
+ if self.allow_degraded:
56
+ logger.warning("%s — bypassing cache", message)
57
+ else:
58
+ raise RuntimeError(message) from error
59
+
60
+ def get_classification(self, text: str) -> dict | None:
61
+ if not self.available:
62
+ return None
63
+ try:
64
+ raw = self._client.get(_text_key(CLASSIFICATION_PREFIX, text))
65
+ return json.loads(raw) if raw else None
66
+ except Exception as error:
67
+ logger.warning("[RedisCache] classification get failed: %s", error)
68
+ return None
69
+
70
+ def set_classification(self, text: str, payload: dict) -> None:
71
+ if not self.available:
72
+ return
73
+ try:
74
+ self._client.setex(
75
+ _text_key(CLASSIFICATION_PREFIX, text),
76
+ self.ttl_seconds,
77
+ json.dumps(payload),
78
+ )
79
+ except Exception as error:
80
+ logger.warning("[RedisCache] classification set failed: %s", error)
81
+
82
+ def get_embedding(self, text: str) -> list[float] | None:
83
+ if not self.available:
84
+ return None
85
+ try:
86
+ raw = self._client.get(_text_key(EMBEDDING_PREFIX, text))
87
+ if not raw:
88
+ return None
89
+ values = json.loads(raw)
90
+ return [float(v) for v in values]
91
+ except Exception as error:
92
+ logger.warning("[RedisCache] embedding get failed: %s", error)
93
+ return None
94
+
95
+ def set_embedding(self, text: str, embedding: list[float]) -> None:
96
+ if not self.available:
97
+ return
98
+ try:
99
+ self._client.setex(
100
+ _text_key(EMBEDDING_PREFIX, text),
101
+ self.ttl_seconds,
102
+ json.dumps(embedding),
103
+ )
104
+ except Exception as error:
105
+ logger.warning("[RedisCache] embedding set failed: %s", error)
106
+
107
+
108
+ redis_cache = RedisInferenceCache()
backend/services/semantic_duplicate_service.py CHANGED
@@ -66,10 +66,18 @@ class SemanticDuplicateService:
66
  Generate a 384-dimensional embedding vector for the given text.
67
  Returns None if the model isn't loaded.
68
  """
 
 
 
 
 
 
69
  if not self.model:
70
  return None
71
  try:
72
- return self.model.encode(text).tolist()
 
 
73
  except Exception as e:
74
  logger.error(f"[SemanticDuplicate] Embedding error: {e}")
75
  return None
 
66
  Generate a 384-dimensional embedding vector for the given text.
67
  Returns None if the model isn't loaded.
68
  """
69
+ from backend.services.redis_cache import redis_cache
70
+
71
+ cached = redis_cache.get_embedding(text)
72
+ if cached is not None:
73
+ return cached
74
+
75
  if not self.model:
76
  return None
77
  try:
78
+ embedding = self.model.encode(text).tolist()
79
+ redis_cache.set_embedding(text, embedding)
80
+ return embedding
81
  except Exception as e:
82
  logger.error(f"[SemanticDuplicate] Embedding error: {e}")
83
  return None
backend/tests/conftest.py CHANGED
@@ -266,7 +266,15 @@ def fake_supabase(fake_db):
266
 
267
 
268
  @pytest.fixture(autouse=True)
269
- def mock_ai_services():
 
 
 
 
 
 
 
 
270
  import backend.main as main
271
 
272
  with patch.object(main.classifier_service, "predict") as mock_v1_predict, \
 
266
 
267
 
268
  @pytest.fixture(autouse=True)
269
+ def mock_ai_services(request):
270
+ if request.node.fspath.basename in {
271
+ "test_redis_cache.py",
272
+ "test_semantic_duplicates.py",
273
+ "test_auth_cookie.py",
274
+ }:
275
+ yield
276
+ return
277
+
278
  import backend.main as main
279
 
280
  with patch.object(main.classifier_service, "predict") as mock_v1_predict, \
backend/tests/test_redis_cache.py ADDED
@@ -0,0 +1,79 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """Tests for Redis inference cache (issue #131)."""
2
+
3
+ from unittest.mock import MagicMock, patch
4
+
5
+ import pytest
6
+
7
+ from backend.services.redis_cache import RedisInferenceCache, _text_key, CLASSIFICATION_PREFIX, EMBEDDING_PREFIX
8
+
9
+
10
+ @pytest.fixture(autouse=True)
11
+ def reset_env(monkeypatch):
12
+ monkeypatch.delenv("USE_REDIS_CACHE", raising=False)
13
+ monkeypatch.delenv("REDIS_URL", raising=False)
14
+ monkeypatch.delenv("ALLOW_DEGRADED_STARTUP", raising=False)
15
+ monkeypatch.delenv("REDIS_CACHE_TTL_SECONDS", raising=False)
16
+
17
+
18
+ def test_text_key_is_stable():
19
+ assert _text_key(CLASSIFICATION_PREFIX, "Hello") == _text_key(CLASSIFICATION_PREFIX, " hello ")
20
+ assert _text_key(CLASSIFICATION_PREFIX, "A") != _text_key(EMBEDDING_PREFIX, "A")
21
+
22
+
23
+ def test_cache_disabled_by_default():
24
+ cache = RedisInferenceCache()
25
+ cache.connect()
26
+ assert cache.available is False
27
+ assert cache.get_classification("ticket") is None
28
+
29
+
30
+ def test_classification_roundtrip(monkeypatch):
31
+ monkeypatch.setenv("USE_REDIS_CACHE", "true")
32
+ client = MagicMock()
33
+ client.ping.return_value = True
34
+ store = {}
35
+
36
+ def setex(key, ttl, value):
37
+ store[key] = value
38
+
39
+ client.get.side_effect = lambda key: store.get(key)
40
+ client.setex.side_effect = setex
41
+
42
+ cache = RedisInferenceCache()
43
+ with patch("redis.from_url", return_value=client):
44
+ cache.connect()
45
+
46
+ payload = {"category": "Billing", "priority": "High"}
47
+ assert cache.get_classification("payment failed") is None
48
+ cache.set_classification("payment failed", payload)
49
+ assert cache.get_classification("payment failed") == payload
50
+
51
+
52
+ def test_embedding_roundtrip(monkeypatch):
53
+ monkeypatch.setenv("USE_REDIS_CACHE", "true")
54
+ client = MagicMock()
55
+ client.ping.return_value = True
56
+ store = {}
57
+
58
+ client.get.side_effect = lambda key: store.get(key)
59
+ client.setex.side_effect = lambda key, ttl, value: store.update({key: value})
60
+
61
+ cache = RedisInferenceCache()
62
+ with patch("redis.from_url", return_value=client):
63
+ cache.connect()
64
+
65
+ vector = [0.1, 0.2, 0.3]
66
+ cache.set_embedding("duplicate ticket", vector)
67
+ assert cache.get_embedding("duplicate ticket") == vector
68
+
69
+
70
+ def test_degraded_startup_bypasses_redis(monkeypatch):
71
+ monkeypatch.setenv("USE_REDIS_CACHE", "true")
72
+ monkeypatch.setenv("ALLOW_DEGRADED_STARTUP", "1")
73
+
74
+ cache = RedisInferenceCache()
75
+ with patch("redis.from_url", side_effect=ConnectionError("down")):
76
+ cache.connect()
77
+
78
+ assert cache.available is False
79
+ assert cache.get_embedding("x") is None
requirements.txt CHANGED
@@ -17,6 +17,7 @@ easyocr
17
  slowapi>=0.1.9
18
  supabase==2.22.4
19
  storage3==2.22.4
 
20
  pytest
21
  pytest-asyncio
22
  httpx
 
17
  slowapi>=0.1.9
18
  supabase==2.22.4
19
  storage3==2.22.4
20
+ redis>=5.0.0
21
  pytest
22
  pytest-asyncio
23
  httpx