167 lines
5.9 KiB
Python
167 lines
5.9 KiB
Python
|
|
import threading
|
||
|
|
from abc import ABC, abstractmethod
|
||
|
|
|
||
|
|
from yuxi.repositories.channel_identity_link_repo import ChannelIdentityLinkRepository
|
||
|
|
|
||
|
|
|
||
|
|
class IdentityLinkStore(ABC):
|
||
|
|
@abstractmethod
|
||
|
|
async def load_all(self) -> dict[str, list[str]]: ...
|
||
|
|
|
||
|
|
@abstractmethod
|
||
|
|
async def add(self, identity: str, channel_type: str, peer_id: str) -> None: ...
|
||
|
|
|
||
|
|
@abstractmethod
|
||
|
|
async def remove(self, channel_type: str, peer_id: str) -> None: ...
|
||
|
|
|
||
|
|
@abstractmethod
|
||
|
|
async def remove_identity(self, identity: str) -> None: ...
|
||
|
|
|
||
|
|
|
||
|
|
class PostgresIdentityLinkStore(IdentityLinkStore):
|
||
|
|
def __init__(self, repo: ChannelIdentityLinkRepository | None = None):
|
||
|
|
self._repo = repo or ChannelIdentityLinkRepository()
|
||
|
|
|
||
|
|
async def load_all(self) -> dict[str, list[str]]:
|
||
|
|
return await self._repo.find_all()
|
||
|
|
|
||
|
|
async def add(self, identity: str, channel_type: str, peer_id: str) -> None:
|
||
|
|
await self._repo.add_link(identity, channel_type, peer_id)
|
||
|
|
|
||
|
|
async def remove(self, channel_type: str, peer_id: str) -> None:
|
||
|
|
await self._repo.remove_link(channel_type, peer_id)
|
||
|
|
|
||
|
|
async def remove_identity(self, identity: str) -> None:
|
||
|
|
await self._repo.remove_identity(identity)
|
||
|
|
|
||
|
|
|
||
|
|
class InMemoryIdentityLinkStore(IdentityLinkStore):
|
||
|
|
def __init__(self):
|
||
|
|
self._data: dict[str, list[str]] = {}
|
||
|
|
|
||
|
|
async def load_all(self) -> dict[str, list[str]]:
|
||
|
|
return dict(self._data)
|
||
|
|
|
||
|
|
async def add(self, identity: str, channel_type: str, peer_id: str) -> None:
|
||
|
|
entry = f"{channel_type}:{peer_id}"
|
||
|
|
existing = self._data.setdefault(identity, [])
|
||
|
|
if entry not in existing:
|
||
|
|
existing.append(entry)
|
||
|
|
|
||
|
|
async def remove(self, channel_type: str, peer_id: str) -> None:
|
||
|
|
entry = f"{channel_type}:{peer_id}"
|
||
|
|
for identity, entries in list(self._data.items()):
|
||
|
|
if entry in entries:
|
||
|
|
entries.remove(entry)
|
||
|
|
if not entries:
|
||
|
|
del self._data[identity]
|
||
|
|
|
||
|
|
async def remove_identity(self, identity: str) -> None:
|
||
|
|
self._data.pop(identity, None)
|
||
|
|
|
||
|
|
|
||
|
|
class IdentityLinkResolver:
|
||
|
|
def __init__(self, store: IdentityLinkStore | None = None):
|
||
|
|
self._store = store or PostgresIdentityLinkStore()
|
||
|
|
self._lock = threading.Lock()
|
||
|
|
self._links: dict[str, list[str]] = {}
|
||
|
|
self._linked_id_to_identity: dict[str, str] = {}
|
||
|
|
self._resolved_cache: dict[str, str] = {}
|
||
|
|
self._loaded = False
|
||
|
|
|
||
|
|
def _rebuild_reverse_index(self) -> None:
|
||
|
|
self._linked_id_to_identity = {}
|
||
|
|
for identity, linked_ids in self._links.items():
|
||
|
|
for linked_id in linked_ids:
|
||
|
|
lid_lower = linked_id.lower()
|
||
|
|
if lid_lower not in self._linked_id_to_identity:
|
||
|
|
self._linked_id_to_identity[lid_lower] = identity
|
||
|
|
|
||
|
|
async def ensure_loaded(self) -> None:
|
||
|
|
if self._loaded:
|
||
|
|
return
|
||
|
|
with self._lock:
|
||
|
|
if self._loaded:
|
||
|
|
return
|
||
|
|
self._loaded = True
|
||
|
|
|
||
|
|
self._links = await self._store.load_all()
|
||
|
|
with self._lock:
|
||
|
|
self._rebuild_reverse_index()
|
||
|
|
self._resolved_cache.clear()
|
||
|
|
|
||
|
|
def set_links(self, links: dict[str, list[str]]) -> None:
|
||
|
|
with self._lock:
|
||
|
|
self._links = links
|
||
|
|
self._resolved_cache.clear()
|
||
|
|
self._rebuild_reverse_index()
|
||
|
|
self._loaded = True
|
||
|
|
|
||
|
|
async def add_link(self, identity: str, channel_type: str, peer_id: str) -> None:
|
||
|
|
await self._store.add(identity, channel_type, peer_id)
|
||
|
|
entry = f"{channel_type}:{peer_id}"
|
||
|
|
with self._lock:
|
||
|
|
existing = self._links.get(identity, [])
|
||
|
|
if entry not in existing:
|
||
|
|
existing.append(entry)
|
||
|
|
self._links[identity] = existing
|
||
|
|
self._resolved_cache.clear()
|
||
|
|
entry_lower = entry.lower()
|
||
|
|
if entry_lower not in self._linked_id_to_identity:
|
||
|
|
self._linked_id_to_identity[entry_lower] = identity
|
||
|
|
|
||
|
|
async def remove_link(self, channel_type: str, peer_id: str) -> None:
|
||
|
|
await self._store.remove(channel_type, peer_id)
|
||
|
|
entry = f"{channel_type}:{peer_id}"
|
||
|
|
with self._lock:
|
||
|
|
for identity, entries in list(self._links.items()):
|
||
|
|
if entry in entries:
|
||
|
|
entries.remove(entry)
|
||
|
|
if not entries:
|
||
|
|
del self._links[identity]
|
||
|
|
entry_lower = entry.lower()
|
||
|
|
self._linked_id_to_identity.pop(entry_lower, None)
|
||
|
|
self._resolved_cache.clear()
|
||
|
|
|
||
|
|
async def remove_identity(self, identity: str) -> None:
|
||
|
|
await self._store.remove_identity(identity)
|
||
|
|
with self._lock:
|
||
|
|
self._links.pop(identity, None)
|
||
|
|
self._resolved_cache.clear()
|
||
|
|
self._rebuild_reverse_index()
|
||
|
|
|
||
|
|
def resolve(self, channel: str, peer_id: str, use_cache: bool = True) -> str:
|
||
|
|
key = f"{channel}:{peer_id}"
|
||
|
|
with self._lock:
|
||
|
|
if use_cache and key in self._resolved_cache:
|
||
|
|
return self._resolved_cache[key]
|
||
|
|
|
||
|
|
if not self._linked_id_to_identity:
|
||
|
|
self._resolved_cache[key] = peer_id
|
||
|
|
return peer_id
|
||
|
|
|
||
|
|
result = self._linked_id_to_identity.get(peer_id.lower())
|
||
|
|
if result is not None:
|
||
|
|
self._resolved_cache[key] = result
|
||
|
|
return result
|
||
|
|
|
||
|
|
result = self._linked_id_to_identity.get(key.lower())
|
||
|
|
if result is not None:
|
||
|
|
self._resolved_cache[key] = result
|
||
|
|
return result
|
||
|
|
|
||
|
|
self._resolved_cache[key] = peer_id
|
||
|
|
return peer_id
|
||
|
|
|
||
|
|
def list_links(self) -> dict[str, list[str]]:
|
||
|
|
with self._lock:
|
||
|
|
return dict(self._links)
|
||
|
|
|
||
|
|
def get_resolved_peer_ids(self) -> set[str]:
|
||
|
|
with self._lock:
|
||
|
|
return set(self._resolved_cache.values())
|
||
|
|
|
||
|
|
def clear_cache(self) -> None:
|
||
|
|
with self._lock:
|
||
|
|
self._resolved_cache.clear()
|