124 lines
4.2 KiB
Python
124 lines
4.2 KiB
Python
from __future__ import annotations
|
|
|
|
import asyncio
|
|
import os
|
|
import time
|
|
from dataclasses import dataclass, field
|
|
|
|
|
|
@dataclass
|
|
class CachedAttachment:
|
|
url: str
|
|
content_type: str
|
|
filename: str | None = None
|
|
local_path: str | None = None
|
|
cached_at: float = field(default_factory=time.time)
|
|
|
|
|
|
class AttachmentCache:
|
|
def __init__(self, ttl: float = 120.0, max_attachments_per_session: int = 50):
|
|
self._storage: dict[str, list[CachedAttachment]] = {}
|
|
self._ttl = ttl
|
|
self._max_attachments_per_session = max_attachments_per_session
|
|
self._lock = asyncio.Lock()
|
|
self._cleanup_task: asyncio.Task | None = None
|
|
self._running = False
|
|
|
|
async def start(self, cleanup_interval: float = 60.0) -> None:
|
|
if self._running:
|
|
return
|
|
self._running = True
|
|
self._cleanup_task = asyncio.create_task(self._cleanup_loop(cleanup_interval))
|
|
|
|
async def stop(self) -> None:
|
|
self._running = False
|
|
if self._cleanup_task:
|
|
self._cleanup_task.cancel()
|
|
try:
|
|
await self._cleanup_task
|
|
except asyncio.CancelledError:
|
|
pass
|
|
self._cleanup_task = None
|
|
|
|
async def _cleanup_loop(self, interval: float) -> None:
|
|
while self._running:
|
|
try:
|
|
await asyncio.sleep(interval)
|
|
await self.cleanup_expired()
|
|
except asyncio.CancelledError:
|
|
break
|
|
|
|
async def add(
|
|
self,
|
|
session_id: str,
|
|
url: str,
|
|
content_type: str,
|
|
filename: str | None = None,
|
|
local_path: str | None = None,
|
|
) -> None:
|
|
async with self._lock:
|
|
if session_id not in self._storage:
|
|
self._storage[session_id] = []
|
|
attachments = self._storage[session_id]
|
|
if len(attachments) >= self._max_attachments_per_session:
|
|
return
|
|
if not any(a.url == url for a in attachments):
|
|
attachments.append(
|
|
CachedAttachment(
|
|
url=url,
|
|
content_type=content_type,
|
|
filename=filename,
|
|
local_path=local_path,
|
|
)
|
|
)
|
|
|
|
async def consume(self, session_id: str) -> list[CachedAttachment]:
|
|
async with self._lock:
|
|
attachments = self._storage.pop(session_id, [])
|
|
now = time.time()
|
|
valid = [a for a in attachments if now - a.cached_at <= self._ttl]
|
|
expired = [a for a in attachments if now - a.cached_at > self._ttl]
|
|
for a in expired:
|
|
await self._remove_local_file(a.local_path)
|
|
return valid
|
|
|
|
async def clear(self, session_id: str) -> None:
|
|
async with self._lock:
|
|
attachments = self._storage.pop(session_id, None)
|
|
if attachments:
|
|
for a in attachments:
|
|
await self._remove_local_file(a.local_path)
|
|
|
|
async def cleanup_expired(self) -> int:
|
|
now = time.time()
|
|
async with self._lock:
|
|
expired_keys: list[str] = []
|
|
for key, attachments in self._storage.items():
|
|
valid = [a for a in attachments if now - a.cached_at <= self._ttl]
|
|
stale = [a for a in attachments if now - a.cached_at > self._ttl]
|
|
self._storage[key] = valid
|
|
for a in stale:
|
|
await self._remove_local_file(a.local_path)
|
|
if not self._storage[key]:
|
|
expired_keys.append(key)
|
|
for key in expired_keys:
|
|
del self._storage[key]
|
|
return len(expired_keys)
|
|
|
|
def stats(self) -> dict:
|
|
total_attachments = 0
|
|
for attachments in self._storage.values():
|
|
total_attachments += len(attachments)
|
|
return {
|
|
"session_count": len(self._storage),
|
|
"total_attachments": total_attachments,
|
|
}
|
|
|
|
@staticmethod
|
|
async def _remove_local_file(path: str | None) -> None:
|
|
if path and os.path.isfile(path):
|
|
try:
|
|
await asyncio.to_thread(os.unlink, path)
|
|
except OSError:
|
|
pass
|