test/app/core/database.py

223 lines
7.4 KiB
Python

# -*- 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()