Files
2026-07-13 12:24:33 +08:00

235 lines
7.8 KiB
Python

# SPDX-License-Identifier: Apache-2.0
"""
Session and SessionManager for tracking per-request state
in the multiprocess cache server.
"""
# Standard
from dataclasses import dataclass, field
from typing import Any, Optional, overload
import threading
import time
# First Party
from lmcache.logging import init_logger
from lmcache.v1.multiprocess.custom_types import IPCCacheServerKey
from lmcache.v1.multiprocess.token_hasher import TokenHasher
from lmcache.v1.periodic_thread import (
PeriodicThread,
PeriodicThreadRegistry,
ThreadLevel,
ThreadRunSummary,
create_periodic_thread,
)
logger = init_logger(__name__)
@dataclass
class Session:
"""Tracks accumulated token IDs and computed chunk hashes for a request.
Thread-safe: all public methods are protected by an internal lock
to allow concurrent access from multiple TP worker threads.
"""
request_id: str
hasher: TokenHasher
token_ids: list[int] = field(default_factory=list)
chunk_hashes: list = field(default_factory=list)
last_prefix_hash: Any = None
num_chunks_processed: int = 0
created_at: float = field(default_factory=time.time)
lookup_ipc_key: Optional[IPCCacheServerKey] = None
extras: dict[str, Any] = field(default_factory=dict)
_lock: threading.Lock = field(default_factory=threading.Lock, repr=False)
def set_tokens(self, full_token_ids: list[int]) -> None:
"""Update the token sequence (idempotent, replaces not extends).
Args:
full_token_ids: Complete token sequence.
"""
with self._lock:
self.token_ids = full_token_ids
@overload
def get_hashes(self, start: int, end: int) -> list: ...
@overload
def get_hashes(self, start: int) -> list: ...
def get_hashes(self, start: int, end: int | None = None) -> list:
"""Compute and return chunk hashes for the [start, end) token range.
Internally computes rolling hashes up to end_chunk, skipping
already-computed chunks.
Two calling conventions are supported (declared via ``@overload``)::
get_hashes(start, end) # explicit end
get_hashes(start) # end = last full-chunk boundary
Args:
start: Start token index (must be aligned to chunk_size).
end: End token index (must be aligned to chunk_size).
When omitted (``None``), automatically set to the last
full-chunk boundary of the current token sequence.
Returns:
List of hash values for chunks in [start_chunk, end_chunk).
"""
chunk_size = self.hasher.chunk_size
assert start % chunk_size == 0, (
f"start ({start}) must be a multiple of chunk_size ({chunk_size})"
)
start_chunk = start // chunk_size
with self._lock:
if end is None:
# No explicit end: use the last full-chunk boundary.
# Lock must be held here because `self.token_ids` may be
# concurrently replaced by `set_tokens` from another thread.
end = len(self.token_ids) - (len(self.token_ids) % chunk_size)
assert end % chunk_size == 0, (
f"end ({end}) must be a multiple of chunk_size ({chunk_size})"
)
end_chunk = end // chunk_size
self._compute_hash(end_chunk)
return self.chunk_hashes[start_chunk:end_chunk]
def _compute_hash(self, end_chunk: int) -> None:
"""Compute rolling hashes up to end_chunk.
Uses cached state to skip already-computed chunks.
Args:
end_chunk: Compute hashes up to (but not including) this chunk.
"""
chunk_size = self.hasher.chunk_size
while self.num_chunks_processed < end_chunk:
cs = self.num_chunks_processed * chunk_size
ce = cs + chunk_size
chunk = self.token_ids[cs:ce]
prefix = (
self.last_prefix_hash
if self.last_prefix_hash is not None
else self.hasher.none_hash
)
h = self.hasher.hash_tokens(chunk, prefix)
self.last_prefix_hash = h
self.chunk_hashes.append(h)
self.num_chunks_processed += 1
class SessionManager:
"""Thread-safe manager for per-request sessions."""
DEFAULT_SESSION_TTL = 600 # 10 minutes
DEFAULT_CLEANUP_INTERVAL = 60.0
def __init__(
self,
hasher: TokenHasher,
ttl: float = DEFAULT_SESSION_TTL,
cleanup_interval: float | None = DEFAULT_CLEANUP_INTERVAL,
) -> None:
self._hasher = hasher
self._ttl = ttl
self._sessions: dict[str, Session] = {}
self._lock = threading.Lock()
self._cleanup_interval = cleanup_interval
self._cleanup_thread: PeriodicThread | None = None
if cleanup_interval is not None and cleanup_interval > 0:
self._cleanup_thread = self._create_cleanup_thread(cleanup_interval)
self._cleanup_thread.start()
def get_or_create(self, request_id: str) -> Session:
"""Get existing session or create a new one.
Args:
request_id: Unique request identifier.
Returns:
The Session for this request_id.
"""
with self._lock:
if request_id not in self._sessions:
self._sessions[request_id] = Session(
request_id=request_id, hasher=self._hasher
)
logger.debug("Created session for request_id=%s", request_id)
return self._sessions[request_id]
def remove(self, request_id: str) -> Optional[Session]:
"""Remove a session by request_id.
Args:
request_id: Unique request identifier.
Returns:
The removed session, or None if no session was found.
"""
with self._lock:
if request_id in self._sessions:
session = self._sessions[request_id]
del self._sessions[request_id]
logger.debug("Removed session for request_id=%s", request_id)
return session
return None
def cleanup_expired(self) -> int:
"""Remove sessions that have exceeded their TTL.
Returns:
Number of sessions removed.
"""
now = time.time()
expired = []
with self._lock:
for rid, session in self._sessions.items():
if now - session.created_at > self._ttl:
expired.append(rid)
for rid in expired:
del self._sessions[rid]
if expired:
logger.info("Cleaned up %d expired sessions", len(expired))
return len(expired)
def active_count(self) -> int:
"""Return the number of active sessions.
Returns:
Number of currently tracked sessions.
"""
with self._lock:
return len(self._sessions)
def close(self) -> None:
"""Stop the background cleanup thread and unregister it."""
if self._cleanup_thread is None:
return
PeriodicThreadRegistry.get_instance().unregister(self._cleanup_thread.name)
self._cleanup_thread.stop()
self._cleanup_thread = None
def _create_cleanup_thread(self, cleanup_interval: float) -> PeriodicThread:
def execute_cleanup() -> ThreadRunSummary:
removed = self.cleanup_expired()
return ThreadRunSummary(
success=True,
message=f"Removed {removed} expired sessions",
extra_info={"removed_sessions": str(removed)},
)
return create_periodic_thread(
name=f"SessionManager-cleanup-thread-{id(self):x}",
interval=cleanup_interval,
execute_fn=execute_cleanup,
level=ThreadLevel.MEDIUM,
)