# -*- coding: utf-8 -*- """PostgreSQL/pgvector storage for registered speaker embeddings.""" from __future__ import annotations import logging from typing import Optional import numpy as np from app.core.config import settings logger = logging.getLogger(__name__) try: import asyncpg except ImportError: # pragma: no cover - runtime dependency is installed in Docker images. asyncpg = None # type: ignore[assignment] class PgSpeakerStorage: """Small asyncpg wrapper shared by speaker registration and ASR matching.""" def __init__(self) -> None: self.pool: Optional[asyncpg.Pool] = None @property def is_enabled(self) -> bool: return bool(settings.SPEAKER_DB_ENABLED) @property def is_connected(self) -> bool: return self.pool is not None async def _create_database_if_not_exists(self) -> None: system_config = { "user": settings.DB_USER, "password": settings.DB_PASSWORD, "database": "postgres", "host": settings.DB_HOST, "port": settings.DB_PORT, } connection: Optional[asyncpg.Connection] = None try: if asyncpg is None: raise RuntimeError("asyncpg is not installed") connection = await asyncpg.connect(**system_config) exists = await connection.fetchval( "SELECT 1 FROM pg_database WHERE datname = $1", settings.DB_NAME, ) if not exists: await connection.execute(f'CREATE DATABASE "{settings.DB_NAME}"') except Exception as exc: logger.warning("尝试自动创建数据库失败,将继续连接业务库: %s", exc) finally: if connection is not None: await connection.close() async def connect(self) -> None: if not self.is_enabled: logger.info("声纹数据库未启用,跳过 PostgreSQL 连接") return if asyncpg is None: raise RuntimeError("asyncpg is not installed") if self.pool is not None: return await self._create_database_if_not_exists() self.pool = await asyncpg.create_pool( user=settings.DB_USER, password=settings.DB_PASSWORD, database=settings.DB_NAME, host=settings.DB_HOST, port=settings.DB_PORT, min_size=1, max_size=settings.DB_POOL_MAX_SIZE, ) async with self.pool.acquire() as connection: await connection.execute("CREATE EXTENSION IF NOT EXISTS vector;") await connection.execute( """ CREATE TABLE IF NOT EXISTS speakers ( id SERIAL PRIMARY KEY, name TEXT NOT NULL UNIQUE, user_id TEXT, embedding vector(192), created_at TIMESTAMP DEFAULT CURRENT_TIMESTAMP ); """ ) await connection.execute( """ DO $$ BEGIN IF NOT EXISTS ( SELECT 1 FROM information_schema.columns WHERE table_name='speakers' AND column_name='user_id' ) THEN ALTER TABLE speakers ADD COLUMN user_id TEXT; END IF; END $$; """ ) await connection.execute( """ DO $$ BEGIN IF NOT EXISTS ( SELECT 1 FROM pg_class c JOIN pg_namespace n ON n.oid = c.relnamespace WHERE c.relname = 'speakers_embedding_idx' ) THEN CREATE INDEX speakers_embedding_idx ON speakers USING hnsw (embedding vector_cosine_ops) WITH (m = 16, ef_construction = 64); END IF; END $$; """ ) logger.info( "声纹数据库已连接: %s:%s/%s", settings.DB_HOST, settings.DB_PORT, settings.DB_NAME, ) async def close(self) -> None: if self.pool is not None: await self.pool.close() self.pool = None def _require_pool(self) -> asyncpg.Pool: if self.pool is None: raise RuntimeError("Speaker database pool is not initialized") return self.pool @staticmethod def _embedding_text(embedding: np.ndarray) -> str: return str(np.asarray(embedding, dtype=np.float32).flatten().tolist()) async def save_speaker( self, name: str, embedding: np.ndarray, user_id: Optional[str] = None, ) -> dict[str, Optional[str]]: pool = self._require_pool() async with pool.acquire() as connection: row = await connection.fetchrow( """ INSERT INTO speakers (name, embedding, user_id) VALUES ($1, $2, $3) ON CONFLICT (name) DO UPDATE SET embedding = EXCLUDED.embedding, user_id = EXCLUDED.user_id RETURNING id, name, user_id; """, name, self._embedding_text(embedding), user_id, ) return { "id": str(row["id"]) if row is not None else None, "name": row["name"] if row is not None else name, "user_id": row["user_id"] if row is not None else user_id, } async def identify_speaker( self, embedding: np.ndarray, threshold: float, ) -> dict[str, Optional[str]]: pool = self._require_pool() distance_threshold = 1.0 - float(threshold) async with pool.acquire() as connection: row = await connection.fetchrow( """ SELECT id, name, user_id, (embedding <=> $1::vector) AS distance FROM speakers ORDER BY embedding <=> $1::vector LIMIT 1; """, self._embedding_text(embedding), ) if row is not None and float(row["distance"]) < distance_threshold: return { "id": str(row["id"]), "name": row["name"], "user_id": row["user_id"], } return {"id": None, "name": None, "user_id": None} async def list_speakers(self) -> list[dict[str, Optional[str]]]: pool = self._require_pool() async with pool.acquire() as connection: rows = await connection.fetch( "SELECT id, name, user_id, created_at FROM speakers ORDER BY id ASC" ) return [ { "id": str(row["id"]), "name": row["name"], "user_id": row["user_id"], "created_at": row["created_at"].isoformat() if row["created_at"] else None, } for row in rows ] async def delete_speaker(self, speaker_id: int) -> bool: pool = self._require_pool() async with pool.acquire() as connection: result = await connection.execute( "DELETE FROM speakers WHERE id = $1", speaker_id, ) return result != "DELETE 0" pg_speaker_db = PgSpeakerStorage()