223 lines
7.4 KiB
Python
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()
|