File size: 3,729 Bytes
00c0691
 
 
 
 
 
 
07d3694
00c0691
 
 
 
 
 
 
 
 
 
 
 
07d3694
00c0691
 
 
 
 
 
 
 
 
 
 
 
 
 
 
880a3c3
 
00c0691
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
b1d2b0b
5063d41
00c0691
 
 
b1d2b0b
 
86f9adc
 
 
 
 
 
5063d41
86f9adc
 
 
 
b1d2b0b
00c0691
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
import time
import uuid
import hashlib
from dataclasses import dataclass, field
from typing import Optional
import chromadb
from chromadb.utils import embedding_functions
from collections import deque
@dataclass
class ChatTurn:
    role: str          # "user" or "assistant"
    content: str
    timestamp: float = field(default_factory=time.time)

@dataclass
class Session:
    session_id: str
    filename: str
    doc_hash: str
    collection: chromadb.Collection
    chat_history: deque[ChatTurn] = field(default_factory=lambda: deque(maxlen=7))
    created_at: float = field(default_factory=time.time)
    last_accessed: float = field(default_factory=time.time)

    def touch(self):
        self.last_accessed = time.time()

    def add_turn(self, role: str, content: str):
        self.chat_history.append(ChatTurn(role, content))
        self.touch()


class SessionManager:
    def __init__(self, max_sessions: int = 10, ttl_seconds: int = 3600):
        self.client = chromadb.EphemeralClient()  # in-memory
        self.embed_fn = embedding_functions.SentenceTransformerEmbeddingFunction(
            model_name="all-MiniLM-L6-v2",
            device="cpu"
        )
        self.sessions: dict[str, Session] = {}
        self.max_sessions = max_sessions
        self.ttl_seconds = ttl_seconds

    def _hash_file(self, file_bytes: bytes) -> str:
        return hashlib.sha256(file_bytes).hexdigest()[:16]

    def create_session(self, filename: str, file_bytes: bytes, chunks: dict[str, list]) -> Session:
        self._evict_if_needed()

        doc_hash = self._hash_file(file_bytes)
        # reuse an existing session if the same doc is already loaded
        for s in self.sessions.values():
            if s.doc_hash == doc_hash:
                s.touch()
                return s

        session_id = str(uuid.uuid4())
        t = time.perf_counter()
        collection = self.client.get_or_create_collection(
            name=f"doc_{session_id}",
            embedding_function=self.embed_fn,
        )
        print(f"    creating collection: {time.perf_counter()-t:.2f}s")
        t = time.perf_counter()
        texts = chunks["text"]
        ids = [f"{session_id}_{i}" for i in range(len(texts))]
        metadatas = chunks["metadata"]
        batch_size = 512
        for start in range(0, len(texts), batch_size):
            end = start + batch_size
            collection.upsert(
                documents=texts[start:end],
                ids=ids[start:end],
                metadatas=metadatas[start:end]
            )
        print(f"  chromadb insert: {time.perf_counter()-t:.2f}s")
        session = Session(
            session_id=session_id,
            filename=filename,
            doc_hash=doc_hash,
            collection=collection,
        )
        self.sessions[session_id] = session
        return session

    def get_session(self, session_id: str) -> Optional[Session]:
        session = self.sessions.get(session_id)
        if session:
            session.touch()
        return session

    def close_session(self, session_id: str):
        session = self.sessions.pop(session_id, None)
        if session:
            self.client.delete_collection(session.collection.name)

    def _evict_if_needed(self):
        now = time.time()
        # TTL eviction first
        expired = [sid for sid, s in self.sessions.items()
                   if now - s.last_accessed > self.ttl_seconds]
        for sid in expired:
            self.close_session(sid)

        # LRU eviction if still over capacity
        if len(self.sessions) >= self.max_sessions:
            oldest = min(self.sessions.values(), key=lambda s: s.last_accessed)
            self.close_session(oldest.session_id)