diff --git a/.env.example b/.env.example index 424755c..f90fcc3 100644 --- a/.env.example +++ b/.env.example @@ -21,6 +21,7 @@ ES_TAG_VECTORS_INDEX=somni_audio_tag_dictionary SOMNI_ES_NODE=http://localhost:9201 SOMNI_ES_AUDIO_INDEX=somni_audio_materials SOMNI_ES_TAG_VECTORS_INDEX=somni_audio_tag_dictionary +SOMNI_ES_SEARCH_EVENTS_INDEX=somni_audio_search_events # MongoDB(功能手板 Fullive) MONGO_URI=mongodb://user:password@host:27017/Fullive @@ -68,10 +69,18 @@ LOG_LEVEL=INFO LOG_DIR=logs LOG_RETENTION=7 days -# 音频检索缓存(空 REDIS_URL 表示关闭;TTL 自写入起算,命中不续期) +# 功能手板 / HTTP 检索缓存(空 REDIS_URL 表示关闭;TTL 自写入起算,命中不续期) REDIS_URL=redis://127.0.0.1:6379/0 -# 量产 Redis(与手板隔离) -SOMNI_REDIS_URL=redis://127.0.0.1:6379/1 +# 量产 Redis(独立实例;空则 GetHot 关闭,不回退 REDIS_URL) +# 本地可另起:redis-server --port 6380 --save "" --appendonly no +SOMNI_REDIS_URL=redis://127.0.0.1:6380/0 +# GetHot 热点排行(Redis ZSET + ES 搜索事件索引) +SOMNI_HOT_ENABLED=true +SOMNI_HOT_TOP_N=10 +SOMNI_HOT_REDIS_KEY=somni:audio:hot:v1 +SOMNI_REDIS_MAX_CONNECTIONS=128 +SOMNI_REDIS_CONNECT_TIMEOUT_SEC=2 +SOMNI_REDIS_SOCKET_TIMEOUT_SEC=2 # 连接池大小:需 ≥ HTTP 并发峰值,过小会 Too many connections REDIS_MAX_CONNECTIONS=512 SEARCH_CACHE_MAX_SIZE=2048 diff --git a/README.md b/README.md index fe9e20e..0bf8aa6 100644 --- a/README.md +++ b/README.md @@ -4,7 +4,7 @@ - **核心**:三维度音频检索(ES 召回 + 标签词典向量 + 精排) - **功能手板**:`MONGO_*` + `ES_NODE` + `REDIS_URL`;HTTP `:8080` + gRPC `:50065` -- **量产**:`SOMNI_MONGO_*` + `SOMNI_ES_*` + `SOMNI_REDIS_URL`;gRPC `:50064` +- **量产**:`SOMNI_MONGO_*` + `SOMNI_ES_*` + `SOMNI_REDIS_URL`(独立 Redis 实例);gRPC `:50064` - **写路径**:直连 Mongo,再同步本侧 ES(不再调用 BioNode) - **接口文档**:`docs/功能手板接口文档.md`、`docs/量产接口文档.md` diff --git a/app/core/config.py b/app/core/config.py index b9568ce..57d87ce 100644 --- a/app/core/config.py +++ b/app/core/config.py @@ -40,6 +40,7 @@ class Settings(BaseSettings): somni_es_node: str = "" somni_es_audio_index: str = "somni_audio_materials" somni_es_tag_vectors_index: str = "somni_audio_tag_dictionary" + somni_es_search_events_index: str = "somni_audio_search_events" mongo_uri: str = "" mongo_db: str = "Fullive" @@ -54,8 +55,8 @@ class Settings(BaseSettings): somni_mongo_answers_collection: str = "somni_quiz_answers" sim_threshold: float = 0.7 # 内容形态向量模糊命中阈值(规范 §五-2) - # GetAudio query_text 与根标签向量相似度下限 - get_audio_root_tag_sim_threshold: float = 0.85 + # GetAudio query_text 与内容形态标签(含二级)向量相似度下限 + get_audio_root_tag_sim_threshold: float = 0.75 # 多路文本检索厌恶硬剔除阈值;≥ 该值 penalty=1.0 丢弃候选 strong_dislike_sim_threshold: float = 0.85 search_sleep_stage_filter_enabled: bool = True # 检索步骤 1 是否按睡眠阶段过滤 @@ -89,8 +90,14 @@ class Settings(BaseSettings): # 功能手板 Redis(空 URL 表示关闭) redis_url: str = "" - # 量产 Redis(与手板隔离) + # 量产 Redis(独立实例;空则 GetHot 关闭,不回退 redis_url) somni_redis_url: str = "" + somni_hot_enabled: bool = True + somni_hot_top_n: int = 10 + somni_hot_redis_key: str = "somni:audio:hot:v1" + somni_redis_max_connections: int = 128 + somni_redis_connect_timeout_sec: float = 2.0 + somni_redis_socket_timeout_sec: float = 2.0 # 连接池需覆盖 HTTP 并发峰值;redis-py 默认仅 100,高并发易 Too many connections redis_max_connections: int = 512 search_cache_max_size: int = 2048 diff --git a/app/core/somni_redis.py b/app/core/somni_redis.py new file mode 100644 index 0000000..7d3475e --- /dev/null +++ b/app/core/somni_redis.py @@ -0,0 +1,40 @@ +"""量产 Redis 客户端(与功能手板/HTTP Redis 物理隔离)。""" + +from __future__ import annotations + +from loguru import logger +from redis.asyncio import Redis + +from app.core.config import Settings + + +def resolve_somni_redis_url(settings: Settings) -> str: + return settings.somni_redis_url.strip() + + +async def create_somni_redis(settings: Settings) -> Redis | None: + """仅按 SOMNI_REDIS_URL 建连接;启用热点时配置错误立即阻止启动。""" + if not settings.somni_hot_enabled: + return None + url = resolve_somni_redis_url(settings) + if not url: + logger.warning("未配置 SOMNI_REDIS_URL,量产 GetHot 热点排行不可用") + return None + client = Redis.from_url( + url, + decode_responses=True, + max_connections=max(1, settings.somni_redis_max_connections), + socket_connect_timeout=max(0.1, settings.somni_redis_connect_timeout_sec), + socket_timeout=max(0.1, settings.somni_redis_socket_timeout_sec), + health_check_interval=30, + ) + try: + await client.ping() + except Exception: + await client.aclose() + raise + logger.info( + "已连接量产独立 Redis,max_connections={}", + settings.somni_redis_max_connections, + ) + return client diff --git a/app/es/search_events.py b/app/es/search_events.py new file mode 100644 index 0000000..4b03793 --- /dev/null +++ b/app/es/search_events.py @@ -0,0 +1,95 @@ +"""量产音频搜索事件 ES 明细。""" + +from __future__ import annotations + +import asyncio +from datetime import UTC, datetime +from typing import Any + +from loguru import logger + +from app.core.config import Settings + +_MAPPING = { + "mappings": { + "properties": { + "keyword": {"type": "keyword"}, + "raw_query": {"type": "keyword"}, + "created_at": {"type": "date"}, + "hit_count": {"type": "integer"}, + "request_id": {"type": "keyword"}, + } + } +} + + +def _build_event_doc( + *, + keyword: str, + raw_query: str, + hit_count: int, + request_id: str, +) -> dict[str, Any]: + return { + "keyword": keyword, + "raw_query": raw_query, + "created_at": datetime.now(UTC).isoformat(), + "hit_count": int(hit_count), + "request_id": request_id, + } + + +class SearchEventsStore: + def __init__(self, client: Any, settings: Settings) -> None: + self._client = client + self._index = settings.somni_es_search_events_index + self._ensure_lock = asyncio.Lock() + self._index_ready = False + + async def ensure_index(self) -> None: + if self._index_ready: + return + async with self._ensure_lock: + if self._index_ready: + return + if await self._client.indices.exists(index=self._index): + self._index_ready = True + return + try: + await self._client.indices.create(index=self._index, body=_MAPPING) + logger.info("已创建 ES 索引:{}", self._index) + except Exception as exc: + if not _is_already_exists_error(exc): + raise + self._index_ready = True + + async def index_event( + self, + *, + keyword: str, + raw_query: str, + hit_count: int, + request_id: str = "", + ) -> None: + await self.ensure_index() + doc = _build_event_doc( + keyword=keyword, + raw_query=raw_query, + hit_count=hit_count, + request_id=request_id, + ) + await self._client.index(index=self._index, document=doc) + + +def _is_already_exists_error(exc: Exception) -> bool: + details = ( + str(exc), + str(getattr(exc, "error", "")), + str(getattr(exc, "body", "")), + str(getattr(exc, "info", "")), + ) + return any( + marker in detail + for detail in details + for marker in ("resource_already_exists_exception", "index_already_exists_exception") + ) diff --git a/app/main.py b/app/main.py index 7607580..b40c2b5 100644 --- a/app/main.py +++ b/app/main.py @@ -55,15 +55,18 @@ def _bootstrap_dev_entry() -> None: from app.core.config import Settings, get_settings from app.core.exception_handlers import register_exception_handlers from app.core.logging import setup_logging +from app.core.somni_redis import create_somni_redis from app.embedding.encoder import Encoder, create_encoder from app.es.client import create_es_client from app.es.search import EsSearch +from app.es.search_events import SearchEventsStore from app.es.sync import EsSync from app.middleware.request_log import register_request_log_middleware from app.server.bootstrap import GrpcServers, start_grpc_servers, stop_grpc_servers from app.server.handboard.audio.service import AudioService from app.server.handboard.audio.store import MaterialsStore, create_materials_store from app.server.somni.audio.catalog import AudioCatalogService as SomniAudioService +from app.server.somni.audio.hot import HotTracker from app.server.somni.quiz.service import QuizService as SomniQuizService from app.server.somni.report.service import ReportService as SomniReportService from app.services.retrieval import RetrievalService @@ -189,11 +192,15 @@ async def lifespan(app: FastAPI) -> AsyncIterator[None]: audio_index=settings.somni_es_audio_index, tag_dictionary_index=settings.somni_es_tag_vectors_index, ) + somni_redis = await create_somni_redis(settings) + events_store = SearchEventsStore(somni_es_client, settings) + hot_tracker = HotTracker(somni_redis, events_store, settings) _app_state.somni_audio_service = SomniAudioService( somni_mongo, settings, es_search=somni_es_search, encoder=encoder, + hot=hot_tracker, ) start_sync_scheduler(_app_state, settings) @@ -206,11 +213,15 @@ async def lifespan(app: FastAPI) -> AsyncIterator[None]: await stop_grpc_servers(_app_state.grpc_servers) _app_state.grpc_servers = None shutdown_sync_scheduler() + if _app_state.somni_audio_service is not None: + await _app_state.somni_audio_service.drain_hot_tasks() await shutdown_audio_search_cache(search_cache) if materials_store is not None: materials_store.close() if somni_mongo is not None: somni_mongo.close() + if somni_redis is not None: + await somni_redis.aclose() await somni_es_client.close() await es_client.close() diff --git a/app/server/somni/audio/catalog.py b/app/server/somni/audio/catalog.py index 16a8963..c266ebb 100644 --- a/app/server/somni/audio/catalog.py +++ b/app/server/somni/audio/catalog.py @@ -7,6 +7,7 @@ from time import monotonic from typing import Any +from loguru import logger from motor.motor_asyncio import AsyncIOMotorClient, AsyncIOMotorCollection from app.core.bson_util import bson_to_jsonable @@ -15,6 +16,7 @@ from app.core.exceptions import AppError, EncoderNotReadyError from app.embedding.encoder import Encoder from app.es.search import EsSearch +from app.server.somni.audio.hot import HotTracker _TAG_ENABLED = "启用" _CONTENT_FORM = "content_form" @@ -33,20 +35,36 @@ def __init__( *, es_search: EsSearch | None = None, encoder: Encoder | None = None, + hot: HotTracker | None = None, ) -> None: self._client = client self._settings = settings self._es_search = es_search self._encoder = encoder + self._hot = hot self._audio_cache: dict[bool, tuple[float, list[dict[str, Any]]]] = {} self._audio_cache_lock = asyncio.Lock() + self._hot_tasks: set[asyncio.Task[None]] = set() + self._hot_sem = asyncio.Semaphore(32) async def get_audio_tag(self) -> dict[str, Any]: collection = self._tags() query = _root_tag_query() total = await collection.count_documents(query) self._reject_over_limit(total) - cursor = collection.find(query, {"type": 1, "code": 1, "name": 1, "name_en": 1}) + cursor = collection.find( + query, + { + "_id": 1, + "id": 1, + "type": 1, + "code": 1, + "name": 1, + "name_en": 1, + "parent_tag_id": 1, + "status": 1, + }, + ) docs = [bson_to_jsonable(doc) async for doc in cursor] return {"tags": [_map_tag_dict(doc) for doc in docs]} @@ -65,12 +83,47 @@ async def get_audio( if code: docs = [doc for doc in docs if _has_content_form_code(doc, code)] if text: - tag_ids = await self._root_tag_ids_by_text(text) - docs = [doc for doc in docs if _has_root_content_form_id(doc, tag_ids)] - return _paginate_docs(docs, page, page_size, fetch_all, self._settings) - - async def get_hot(self) -> None: - return None + tag_ids = await self._content_form_tag_ids_by_text(text) + docs = [doc for doc in docs if _has_content_form_tag_id(doc, tag_ids)] + payload = _paginate_docs(docs, page, page_size, fetch_all, self._settings) + payload["list"] = [_to_audio_list_item(item) for item in payload["list"]] + self._schedule_hot(query_text, int(payload.get("total") or 0)) + return payload + + async def get_hot(self) -> dict[str, Any]: + if self._hot is None: + raise AppError( + message="量产 Redis 未配置,无法获取热点", + status_code=HttpStatus.SERVICE_UNAVAILABLE, + ) + return {"items": await self._hot.list_hot()} + + def _schedule_hot(self, query_text: str, hit_count: int) -> None: + if self._hot is None or not query_text.strip(): + return + task = asyncio.create_task(self._record_hot_safely(query_text, hit_count)) + self._hot_tasks.add(task) + task.add_done_callback(self._hot_tasks.discard) + + async def drain_hot_tasks(self, *, timeout_sec: float = 5.0) -> None: + """关闭前排空热点记账任务,避免访问已关闭的 Redis/ES。""" + pending = [task for task in self._hot_tasks if not task.done()] + if not pending: + return + done, still = await asyncio.wait(pending, timeout=max(0.1, timeout_sec)) + for task in still: + task.cancel() + if still: + await asyncio.gather(*still, return_exceptions=True) + logger.warning("量产热点后台任务关闭超时,已取消 {} 个", len(still)) + _ = done + + async def _record_hot_safely(self, query_text: str, hit_count: int) -> None: + async with self._hot_sem: + try: + await self._hot.record_search(query_text, hit_count=hit_count) + except Exception as exc: + logger.warning("量产热点后台记账失败:{}", exc) async def _load_audios(self, *, from_es: bool) -> list[dict[str, Any]]: now = monotonic() @@ -109,7 +162,7 @@ async def _fetch_audios_es(self) -> list[dict[str, Any]]: self._reject_over_limit(len(docs)) return docs - async def _root_tag_ids_by_text(self, text: str) -> set[str]: + async def _content_form_tag_ids_by_text(self, text: str) -> set[str]: if self._encoder is None or not self._encoder.is_loaded: raise EncoderNotReadyError() if self._es_search is None: @@ -120,19 +173,20 @@ async def _root_tag_ids_by_text(self, text: str) -> set[str]: query_vector = await self._encoder.encode_one(text) tags = await self._es_search.list_content_tag_vectors() threshold = self._settings.get_audio_root_tag_sim_threshold - matched: set[str] = set() + scored: list[tuple[float, str]] = [] for tag in tags: - if not _is_root_content_form_dict(tag): + if not _is_content_form_dict(tag): continue vector = tag.get("vector") if not isinstance(vector, list) or not vector: continue - if _cosine_similarity(query_vector, vector) <= threshold: + sim = _cosine_similarity(query_vector, vector) + if sim <= threshold: continue tag_id = str(tag.get("id") or "").strip() if tag_id: - matched.add(tag_id) - return matched + scored.append((sim, tag_id)) + return _select_matched_tag_ids(scored) def _tags(self) -> AsyncIOMotorCollection: return self._db()[self._settings.somni_mongo_tag_dictionary_collection] @@ -154,8 +208,19 @@ def _reject_over_limit(self, total: int) -> None: raise InvalidAudioQueryError(f"全量条数超过上限 {limit}") +def _select_matched_tag_ids(scored: list[tuple[float, str]]) -> set[str]: + """保留高分标签;近精确命中时收紧范围,避免宽泛根标签稀释结果。""" + if not scored: + return set() + best = max(sim for sim, _ in scored) + if best >= 0.9: + return {tag_id for sim, tag_id in scored if sim >= best - 0.05} + return {tag_id for _, tag_id in scored} + + def _root_tag_query() -> dict[str, Any]: return { + "type": _CONTENT_FORM, "status": _TAG_ENABLED, "$or": [ {"parent_tag_id": {"$exists": False}}, @@ -176,12 +241,11 @@ def _paginate_docs( if fetch_all: if total > settings.fetch_all_hard_limit: raise InvalidAudioQueryError(f"全量条数超过上限 {settings.fetch_all_hard_limit}") - return {"materials": docs, "page": _page_info(1, len(docs), total, 1)} + return {"list": docs, "page": 1, "page_size": len(docs), "total": total} cur_page, size = _page_window(page, page_size, settings) start = (cur_page - 1) * size chunk = docs[start : start + size] - pages = math.ceil(total / size) if size else 0 - return {"materials": chunk, "page": _page_info(cur_page, size, total, pages)} + return {"list": chunk, "page": cur_page, "page_size": size, "total": total} def _page_window( @@ -196,34 +260,70 @@ def _page_window( return cur_page, min(size, settings.max_page_size) -def _page_info(page: int, page_size: int, total: int, total_pages: int) -> dict[str, int]: - return {"page": page, "page_size": page_size, "total": total, "total_pages": total_pages} - - -def _map_tag_dict(doc: dict[str, Any]) -> dict[str, str]: +def _map_tag_dict(doc: dict[str, Any]) -> dict[str, Any]: + parent = doc.get("parent_tag_id") return { "type": str(doc.get("type") or ""), "code": str(doc.get("code") or ""), "name": str(doc.get("name") or ""), "name_en": str(doc.get("name_en") or ""), + "id": str(doc.get("id") or doc.get("_id") or ""), + "parent_tag_id": None if parent is None else str(parent), + "status": str(doc.get("status") or ""), } def _map_material(doc: dict[str, Any]) -> dict[str, Any]: - mapped = dict(doc) - mapped.pop("embedding", None) - mapped["id"] = str(doc.get("id") or doc.get("_id") or "") - mapped.pop("_id", None) - return mapped + """缓存/过滤用中间形态,保留 content_form_tags。""" + return { + "id": str(doc.get("id") or doc.get("_id") or ""), + "audio_name": str(doc.get("audio_name") or ""), + "audio_url": str(doc.get("audio_url") or ""), + "cover_url": str(doc.get("cover_url") or ""), + "description": str(doc.get("description") or ""), + "vip": _to_vip(doc.get("vip")), + "content_form_tags": doc.get("content_form_tags") or [], + } + + +def _to_audio_list_item(doc: dict[str, Any]) -> dict[str, Any]: + return { + "id": str(doc.get("id") or ""), + "audio_name": str(doc.get("audio_name") or ""), + "audio_url": str(doc.get("audio_url") or ""), + "cover_url": str(doc.get("cover_url") or ""), + "description": str(doc.get("description") or ""), + "vip": _to_vip(doc.get("vip")), + } + + +def _to_vip(value: Any) -> int: + """库无 vip / 假值时返回 0;真值返回 1(兼容 bool/int/常见字符串)。""" + if value is None or value is False: + return 0 + if isinstance(value, (int, float)): + return 1 if value != 0 else 0 + if isinstance(value, str): + normalized = value.strip().lower() + if normalized in {"", "0", "false", "no", "off", "none", "null"}: + return 0 + if normalized in {"1", "true", "yes", "on"}: + return 1 + return 0 + return 1 if bool(value) else 0 def _is_blank(value: Any) -> bool: return value is None or str(value).strip() in ("", "None") -def _is_root_content_form_dict(tag: dict[str, Any]) -> bool: +def _is_content_form_dict(tag: dict[str, Any]) -> bool: dimension = str(tag.get("dimension") or tag.get("type") or "") - return dimension == _CONTENT_FORM and _is_blank(tag.get("parent_tag_id")) + return dimension == _CONTENT_FORM + + +def _is_root_content_form_dict(tag: dict[str, Any]) -> bool: + return _is_content_form_dict(tag) and _is_blank(tag.get("parent_tag_id")) def _has_content_form_code(doc: dict[str, Any], tag_code: str) -> bool: @@ -233,11 +333,11 @@ def _has_content_form_code(doc: dict[str, Any], tag_code: str) -> bool: return False -def _has_root_content_form_id(doc: dict[str, Any], tag_ids: set[str]) -> bool: +def _has_content_form_tag_id(doc: dict[str, Any], tag_ids: set[str]) -> bool: if not tag_ids: return False for item in doc.get("content_form_tags") or []: - if not isinstance(item, dict) or not _is_blank(item.get("parent_tag_id")): + if not isinstance(item, dict): continue if str(item.get("tag_id") or "") in tag_ids: return True diff --git a/app/server/somni/audio/hot.py b/app/server/somni/audio/hot.py new file mode 100644 index 0000000..c9fba83 --- /dev/null +++ b/app/server/somni/audio/hot.py @@ -0,0 +1,93 @@ +"""量产音频搜索热点:Redis 计数 + ES 明细。""" + +from __future__ import annotations + +from typing import Any + +from loguru import logger + +from app.core.codes import HttpStatus +from app.core.config import Settings +from app.core.exceptions import AppError +from app.es.search_events import SearchEventsStore + + +def normalize_keyword(text: str) -> str: + return text.strip() + + +def _as_str(member: bytes | str) -> str: + if isinstance(member, bytes): + return member.decode() + return member + + +class HotTracker: + def __init__( + self, + redis: Any, + events_store: SearchEventsStore | None, + settings: Settings, + ) -> None: + self._redis = redis + self._events = events_store + self._settings = settings + + async def record_search(self, raw_query: str, *, hit_count: int) -> None: + if not self._settings.somni_hot_enabled: + return + keyword = normalize_keyword(raw_query) + if not keyword: + return + await self._safe_redis_incr(keyword) + await self._safe_es_index(keyword, raw_query, hit_count) + + async def list_hot(self) -> list[dict[str, Any]]: + if not self._settings.somni_hot_enabled: + return [] + if self._settings.somni_hot_top_n <= 0: + return [] + if self._redis is None: + raise AppError( + message="量产 Redis 未配置,无法获取热点", + status_code=HttpStatus.SERVICE_UNAVAILABLE, + ) + try: + rows = await self._redis.zrevrange( + self._settings.somni_hot_redis_key, + 0, + self._settings.somni_hot_top_n - 1, + withscores=True, + ) + except Exception as exc: + logger.warning("量产热点 Redis 读取失败:{}", exc) + raise AppError( + message="量产 Redis 不可用,无法获取热点", + status_code=HttpStatus.SERVICE_UNAVAILABLE, + ) from exc + return [{"keyword": _as_str(member), "score": int(score)} for member, score in rows] + + async def _safe_redis_incr(self, keyword: str) -> None: + if self._redis is None: + return + try: + await self._redis.zincrby(self._settings.somni_hot_redis_key, 1, keyword) + except Exception as exc: + logger.warning("量产热点 Redis 写入失败:{}", exc) + + async def _safe_es_index( + self, + keyword: str, + raw_query: str, + hit_count: int, + ) -> None: + if self._events is None: + return + try: + await self._events.index_event( + keyword=keyword, + raw_query=raw_query, + hit_count=hit_count, + ) + except Exception as exc: + logger.warning("量产热点 ES 写入失败:{}", exc) diff --git a/app/server/somni/audio/rpc.py b/app/server/somni/audio/rpc.py index b0a4d94..849c4ab 100644 --- a/app/server/somni/audio/rpc.py +++ b/app/server/somni/audio/rpc.py @@ -4,7 +4,7 @@ from typing import TYPE_CHECKING, Any -from google.protobuf.struct_pb2 import Struct +from google.protobuf.struct_pb2 import Value from app.core.exceptions import ServiceNotReadyError from app.server.errors import abort_from_app_error, run_rpc_call @@ -42,8 +42,8 @@ async def GetHot(self, request, context): service = await self._require(context) async def _do(): - await service.get_hot() - return uburnode_somni_pb2.GetHotRes() + payload = await service.get_hot() + return _to_hot_res(payload) return await run_rpc_call(context, _do) @@ -72,24 +72,50 @@ def _to_tag_res(payload: dict[str, Any]) -> uburnode_somni_pb2.GetAudioTagRes: code=str(item.get("code") or ""), name=str(item.get("name") or ""), name_en=str(item.get("name_en") or ""), + id=str(item.get("id") or ""), + parent_tag_id=_to_value(item.get("parent_tag_id")), + status=str(item.get("status") or ""), ) ) return res -def _to_audio_res(payload: dict[str, Any]) -> uburnode_somni_pb2.GetAudioRes: - res = uburnode_somni_pb2.GetAudioRes() - for item in payload.get("materials") or []: - struct = Struct() - struct.update(item if isinstance(item, dict) else {}) - res.materials.append(struct) - page = payload.get("page") or {} - res.page.CopyFrom( - uburnode_somni_pb2.PageInfo( - page=int(page.get("page") or 1), - page_size=int(page.get("page_size") or 0), - total=int(page.get("total") or 0), - total_pages=int(page.get("total_pages") or 0), +def _to_value(value: Any) -> Value: + result = Value() + if value is None: + result.null_value = 0 + else: + result.string_value = str(value) + return result + + +def _to_hot_res(payload: dict[str, Any]) -> uburnode_somni_pb2.GetHotRes: + res = uburnode_somni_pb2.GetHotRes() + for item in payload.get("items") or []: + res.items.append( + uburnode_somni_pb2.HotKeyword( + keyword=str(item.get("keyword") or ""), + score=int(item.get("score") or 0), + ) ) + return res + + +def _to_audio_res(payload: dict[str, Any]) -> uburnode_somni_pb2.GetAudioRes: + res = uburnode_somni_pb2.GetAudioRes( + page=int(payload.get("page") or 1), + page_size=int(payload.get("page_size") or 0), + total=int(payload.get("total") or 0), ) + for item in payload.get("list") or []: + res.list.append( + uburnode_somni_pb2.AudioListItem( + id=str(item.get("id") or ""), + audio_name=str(item.get("audio_name") or ""), + audio_url=str(item.get("audio_url") or ""), + cover_url=str(item.get("cover_url") or ""), + description=str(item.get("description") or ""), + vip=int(item.get("vip") or 0), + ) + ) return res diff --git a/app/uburnode_grpc/grpc_gen/uburnode_somni_pb2.py b/app/uburnode_grpc/grpc_gen/uburnode_somni_pb2.py index e9c8ae2..75b958f 100644 --- a/app/uburnode_grpc/grpc_gen/uburnode_somni_pb2.py +++ b/app/uburnode_grpc/grpc_gen/uburnode_somni_pb2.py @@ -25,7 +25,7 @@ from google.protobuf import struct_pb2 as google_dot_protobuf_dot_struct__pb2 -DESCRIPTOR = _descriptor_pool.Default().AddSerializedFile(b'\n\x14uburnode_somni.proto\x12\x11uburnode.somni.v1\x1a\x1cgoogle/protobuf/struct.proto\".\n\x0cGetAnswerReq\x12\x0b\n\x03uid\x18\x01 \x01(\t\x12\x11\n\tanswer_id\x18\x02 \x01(\t\"\x1f\n\x0cGetAnswerRes\x12\x0f\n\x07\x61nswers\x18\x01 \x01(\t\"1\n\rReportDateReq\x12\x0b\n\x03uid\x18\x01 \x01(\t\x12\x13\n\x0brecord_date\x18\x02 \x01(\t\"\x0f\n\rGetSummaryRes\"\x0e\n\x0cGetEventsRes\"\x13\n\x11GetEnvironmentRes\"\x11\n\x0fGetStructureRes\"\x14\n\x12GetSleepQualityRes\"\xc1\x01\n\x0bGetAudioReq\x12\x11\n\x04page\x18\x01 \x01(\x05H\x00\x88\x01\x01\x12\x16\n\tpage_size\x18\x02 \x01(\x05H\x01\x88\x01\x01\x12\x16\n\tfetch_all\x18\x03 \x01(\x08H\x02\x88\x01\x01\x12\x17\n\nquery_text\x18\x04 \x01(\tH\x03\x88\x01\x01\x12\x15\n\x08tag_code\x18\x05 \x01(\tH\x04\x88\x01\x01\x42\x07\n\x05_pageB\x0c\n\n_page_sizeB\x0c\n\n_fetch_allB\r\n\x0b_query_textB\x0b\n\t_tag_code\"O\n\x08PageInfo\x12\x0c\n\x04page\x18\x01 \x01(\x05\x12\x11\n\tpage_size\x18\x02 \x01(\x05\x12\r\n\x05total\x18\x03 \x01(\x05\x12\x13\n\x0btotal_pages\x18\x04 \x01(\x05\"d\n\x0bGetAudioRes\x12*\n\tmaterials\x18\x01 \x03(\x0b\x32\x17.google.protobuf.Struct\x12)\n\x04page\x18\x02 \x01(\x0b\x32\x1b.uburnode.somni.v1.PageInfo\"\x10\n\x0eGetAudioTagReq\"H\n\x0bTagDictItem\x12\x0c\n\x04type\x18\x01 \x01(\t\x12\x0c\n\x04\x63ode\x18\x02 \x01(\t\x12\x0c\n\x04name\x18\x03 \x01(\t\x12\x0f\n\x07name_en\x18\x04 \x01(\t\">\n\x0eGetAudioTagRes\x12,\n\x04tags\x18\x01 \x03(\x0b\x32\x1e.uburnode.somni.v1.TagDictItem\"\x0b\n\tGetHotReq\"\x0b\n\tGetHotRes2\\\n\x0bQuizService\x12M\n\tGetAnswer\x12\x1f.uburnode.somni.v1.GetAnswerReq\x1a\x1f.uburnode.somni.v1.GetAnswerRes2\xbd\x03\n\rReportService\x12P\n\nGetSummary\x12 .uburnode.somni.v1.ReportDateReq\x1a .uburnode.somni.v1.GetSummaryRes\x12N\n\tGetEvents\x12 .uburnode.somni.v1.ReportDateReq\x1a\x1f.uburnode.somni.v1.GetEventsRes\x12X\n\x0eGetEnvironment\x12 .uburnode.somni.v1.ReportDateReq\x1a$.uburnode.somni.v1.GetEnvironmentRes\x12T\n\x0cGetStructure\x12 .uburnode.somni.v1.ReportDateReq\x1a\".uburnode.somni.v1.GetStructureRes\x12Z\n\x0fGetSleepQuality\x12 .uburnode.somni.v1.ReportDateReq\x1a%.uburnode.somni.v1.GetSleepQualityRes2\xf5\x01\n\x0c\x41udioService\x12J\n\x08GetAudio\x12\x1e.uburnode.somni.v1.GetAudioReq\x1a\x1e.uburnode.somni.v1.GetAudioRes\x12S\n\x0bGetAudioTag\x12!.uburnode.somni.v1.GetAudioTagReq\x1a!.uburnode.somni.v1.GetAudioTagRes\x12\x44\n\x06GetHot\x12\x1c.uburnode.somni.v1.GetHotReq\x1a\x1c.uburnode.somni.v1.GetHotResb\x06proto3') +DESCRIPTOR = _descriptor_pool.Default().AddSerializedFile(b'\n\x14uburnode_somni.proto\x12\x11uburnode.somni.v1\x1a\x1cgoogle/protobuf/struct.proto\".\n\x0cGetAnswerReq\x12\x0b\n\x03uid\x18\x01 \x01(\t\x12\x11\n\tanswer_id\x18\x02 \x01(\t\"\x1f\n\x0cGetAnswerRes\x12\x0f\n\x07\x61nswers\x18\x01 \x01(\t\"1\n\rReportDateReq\x12\x0b\n\x03uid\x18\x01 \x01(\t\x12\x13\n\x0brecord_date\x18\x02 \x01(\t\"\x0f\n\rGetSummaryRes\"\x0e\n\x0cGetEventsRes\"\x13\n\x11GetEnvironmentRes\"\x11\n\x0fGetStructureRes\"\x14\n\x12GetSleepQualityRes\"\xc1\x01\n\x0bGetAudioReq\x12\x11\n\x04page\x18\x01 \x01(\x05H\x00\x88\x01\x01\x12\x16\n\tpage_size\x18\x02 \x01(\x05H\x01\x88\x01\x01\x12\x16\n\tfetch_all\x18\x03 \x01(\x08H\x02\x88\x01\x01\x12\x17\n\nquery_text\x18\x04 \x01(\tH\x03\x88\x01\x01\x12\x15\n\x08tag_code\x18\x05 \x01(\tH\x04\x88\x01\x01\x42\x07\n\x05_pageB\x0c\n\n_page_sizeB\x0c\n\n_fetch_allB\r\n\x0b_query_textB\x0b\n\t_tag_code\"\xdf\x01\n\rAudioListItem\x12\x0f\n\x02id\x18\x01 \x01(\tH\x00\x88\x01\x01\x12\x17\n\naudio_name\x18\x02 \x01(\tH\x01\x88\x01\x01\x12\x16\n\taudio_url\x18\x03 \x01(\tH\x02\x88\x01\x01\x12\x16\n\tcover_url\x18\x04 \x01(\tH\x03\x88\x01\x01\x12\x18\n\x0b\x64\x65scription\x18\x05 \x01(\tH\x04\x88\x01\x01\x12\x10\n\x03vip\x18\x06 \x01(\x05H\x05\x88\x01\x01\x42\x05\n\x03_idB\r\n\x0b_audio_nameB\x0c\n\n_audio_urlB\x0c\n\n_cover_urlB\x0e\n\x0c_descriptionB\x06\n\x04_vip\"m\n\x0bGetAudioRes\x12.\n\x04list\x18\x01 \x03(\x0b\x32 .uburnode.somni.v1.AudioListItem\x12\x0c\n\x04page\x18\x02 \x01(\x05\x12\x11\n\tpage_size\x18\x03 \x01(\x05\x12\r\n\x05total\x18\x04 \x01(\x05\"\x10\n\x0eGetAudioTagReq\"\x93\x01\n\x0bTagDictItem\x12\x0c\n\x04type\x18\x01 \x01(\t\x12\x0c\n\x04\x63ode\x18\x02 \x01(\t\x12\x0c\n\x04name\x18\x03 \x01(\t\x12\x0f\n\x07name_en\x18\x04 \x01(\t\x12\n\n\x02id\x18\x05 \x01(\t\x12-\n\rparent_tag_id\x18\x06 \x01(\x0b\x32\x16.google.protobuf.Value\x12\x0e\n\x06status\x18\x07 \x01(\t\">\n\x0eGetAudioTagRes\x12,\n\x04tags\x18\x01 \x03(\x0b\x32\x1e.uburnode.somni.v1.TagDictItem\"\x0b\n\tGetHotReq\",\n\nHotKeyword\x12\x0f\n\x07keyword\x18\x01 \x01(\t\x12\r\n\x05score\x18\x02 \x01(\x03\"9\n\tGetHotRes\x12,\n\x05items\x18\x01 \x03(\x0b\x32\x1d.uburnode.somni.v1.HotKeyword2\\\n\x0bQuizService\x12M\n\tGetAnswer\x12\x1f.uburnode.somni.v1.GetAnswerReq\x1a\x1f.uburnode.somni.v1.GetAnswerRes2\xbd\x03\n\rReportService\x12P\n\nGetSummary\x12 .uburnode.somni.v1.ReportDateReq\x1a .uburnode.somni.v1.GetSummaryRes\x12N\n\tGetEvents\x12 .uburnode.somni.v1.ReportDateReq\x1a\x1f.uburnode.somni.v1.GetEventsRes\x12X\n\x0eGetEnvironment\x12 .uburnode.somni.v1.ReportDateReq\x1a$.uburnode.somni.v1.GetEnvironmentRes\x12T\n\x0cGetStructure\x12 .uburnode.somni.v1.ReportDateReq\x1a\".uburnode.somni.v1.GetStructureRes\x12Z\n\x0fGetSleepQuality\x12 .uburnode.somni.v1.ReportDateReq\x1a%.uburnode.somni.v1.GetSleepQualityRes2\xf5\x01\n\x0c\x41udioService\x12J\n\x08GetAudio\x12\x1e.uburnode.somni.v1.GetAudioReq\x1a\x1e.uburnode.somni.v1.GetAudioRes\x12S\n\x0bGetAudioTag\x12!.uburnode.somni.v1.GetAudioTagReq\x1a!.uburnode.somni.v1.GetAudioTagRes\x12\x44\n\x06GetHot\x12\x1c.uburnode.somni.v1.GetHotReq\x1a\x1c.uburnode.somni.v1.GetHotResb\x06proto3') _globals = globals() _builder.BuildMessageAndEnumDescriptors(DESCRIPTOR, _globals) @@ -50,24 +50,26 @@ _globals['_GETSLEEPQUALITYRES']._serialized_end=298 _globals['_GETAUDIOREQ']._serialized_start=301 _globals['_GETAUDIOREQ']._serialized_end=494 - _globals['_PAGEINFO']._serialized_start=496 - _globals['_PAGEINFO']._serialized_end=575 - _globals['_GETAUDIORES']._serialized_start=577 - _globals['_GETAUDIORES']._serialized_end=677 - _globals['_GETAUDIOTAGREQ']._serialized_start=679 - _globals['_GETAUDIOTAGREQ']._serialized_end=695 - _globals['_TAGDICTITEM']._serialized_start=697 - _globals['_TAGDICTITEM']._serialized_end=769 - _globals['_GETAUDIOTAGRES']._serialized_start=771 - _globals['_GETAUDIOTAGRES']._serialized_end=833 - _globals['_GETHOTREQ']._serialized_start=835 - _globals['_GETHOTREQ']._serialized_end=846 - _globals['_GETHOTRES']._serialized_start=848 - _globals['_GETHOTRES']._serialized_end=859 - _globals['_QUIZSERVICE']._serialized_start=861 - _globals['_QUIZSERVICE']._serialized_end=953 - _globals['_REPORTSERVICE']._serialized_start=956 - _globals['_REPORTSERVICE']._serialized_end=1401 - _globals['_AUDIOSERVICE']._serialized_start=1404 - _globals['_AUDIOSERVICE']._serialized_end=1649 + _globals['_AUDIOLISTITEM']._serialized_start=497 + _globals['_AUDIOLISTITEM']._serialized_end=720 + _globals['_GETAUDIORES']._serialized_start=722 + _globals['_GETAUDIORES']._serialized_end=831 + _globals['_GETAUDIOTAGREQ']._serialized_start=833 + _globals['_GETAUDIOTAGREQ']._serialized_end=849 + _globals['_TAGDICTITEM']._serialized_start=852 + _globals['_TAGDICTITEM']._serialized_end=999 + _globals['_GETAUDIOTAGRES']._serialized_start=1001 + _globals['_GETAUDIOTAGRES']._serialized_end=1063 + _globals['_GETHOTREQ']._serialized_start=1065 + _globals['_GETHOTREQ']._serialized_end=1076 + _globals['_HOTKEYWORD']._serialized_start=1078 + _globals['_HOTKEYWORD']._serialized_end=1122 + _globals['_GETHOTRES']._serialized_start=1124 + _globals['_GETHOTRES']._serialized_end=1181 + _globals['_QUIZSERVICE']._serialized_start=1183 + _globals['_QUIZSERVICE']._serialized_end=1275 + _globals['_REPORTSERVICE']._serialized_start=1278 + _globals['_REPORTSERVICE']._serialized_end=1723 + _globals['_AUDIOSERVICE']._serialized_start=1726 + _globals['_AUDIOSERVICE']._serialized_end=1971 # @@protoc_insertion_point(module_scope) diff --git a/docker-compose.yml b/docker-compose.yml index 28356d0..2b6d346 100644 --- a/docker-compose.yml +++ b/docker-compose.yml @@ -27,6 +27,7 @@ services: environment: ES_NODE: http://elasticsearch:9200 REDIS_URL: redis://redis:6379/0 + SOMNI_REDIS_URL: redis://redis-somni:6379/0 ports: - "50065:50065" # 功能手板 gRPC - "50064:50064" # 量产 gRPC @@ -43,6 +44,8 @@ services: condition: service_healthy redis: condition: service_healthy + redis-somni: + condition: service_healthy restart: unless-stopped redis: @@ -56,6 +59,19 @@ services: retries: 10 restart: unless-stopped + redis-somni: + image: redis:7.4-alpine + container_name: uburnode-redis-somni + command: ["redis-server", "--appendonly", "yes", "--appendfsync", "everysec"] + volumes: + - redis_somni_data:/data + healthcheck: + test: ["CMD", "redis-cli", "ping"] + interval: 5s + timeout: 3s + retries: 10 + restart: unless-stopped + elasticsearch: image: docker.elastic.co/elasticsearch/elasticsearch:8.15.0 container_name: uburnode-es @@ -80,3 +96,4 @@ services: volumes: es_data: hf_cache: + redis_somni_data: diff --git a/proto/uburnode_somni.proto b/proto/uburnode_somni.proto index b2fa765..d6ba398 100644 --- a/proto/uburnode_somni.proto +++ b/proto/uburnode_somni.proto @@ -51,16 +51,20 @@ message GetAudioReq { optional string tag_code = 5; // 内容形态标签 code;空则不按标签过滤 } -message PageInfo { - int32 page = 1; - int32 page_size = 2; - int32 total = 3; - int32 total_pages = 4; +message AudioListItem { + optional string id = 1; + optional string audio_name = 2; + optional string audio_url = 3; + optional string cover_url = 4; + optional string description = 5; + optional int32 vip = 6; // 库无该字段时显式返回 0 } message GetAudioRes { - repeated google.protobuf.Struct materials = 1; - PageInfo page = 2; + repeated AudioListItem list = 1; + int32 page = 2; + int32 page_size = 3; + int32 total = 4; } message GetAudioTagReq {} @@ -70,6 +74,9 @@ message TagDictItem { string code = 2; string name = 3; string name_en = 4; + string id = 5; + google.protobuf.Value parent_tag_id = 6; // 库为 null/缺省时显式返回 null + string status = 7; } message GetAudioTagRes { @@ -78,7 +85,14 @@ message GetAudioTagRes { message GetHotReq {} -message GetHotRes {} +message HotKeyword { + string keyword = 1; + int64 score = 2; +} + +message GetHotRes { + repeated HotKeyword items = 1; +} service AudioService { rpc GetAudio (GetAudioReq) returns (GetAudioRes); diff --git a/tests/test_grpc_somni_audio.py b/tests/test_grpc_somni_audio.py index 48486c9..c58a60a 100644 --- a/tests/test_grpc_somni_audio.py +++ b/tests/test_grpc_somni_audio.py @@ -27,8 +27,19 @@ async def test_get_audio_passes_page_and_query() -> None: rpc, service = _make_rpc() service.get_audio = AsyncMock( return_value={ - "materials": [{"id": "m1", "audio_name": "雨声"}], - "page": {"page": 1, "page_size": 20, "total": 1, "total_pages": 1}, + "list": [ + { + "id": "m1", + "audio_name": "雨声", + "audio_url": "https://cdn.example/a.mp3", + "cover_url": "https://cdn.example/a.png", + "description": "desc", + "vip": 0, + } + ], + "page": 1, + "page_size": 20, + "total": 1, } ) req = uburnode_somni_pb2.GetAudioReq(page=1, page_size=20, query_text="雨声") @@ -40,8 +51,11 @@ async def test_get_audio_passes_page_and_query() -> None: query_text="雨声", tag_code="", ) - assert res.materials[0]["audio_name"] == "雨声" - assert res.page.total == 1 + assert res.list[0].audio_name == "雨声" + assert res.list[0].vip == 0 + assert res.total == 1 + assert res.page == 1 + assert res.page_size == 20 @pytest.mark.asyncio @@ -49,8 +63,10 @@ async def test_get_audio_fetch_all() -> None: rpc, service = _make_rpc() service.get_audio = AsyncMock( return_value={ - "materials": [], - "page": {"page": 1, "page_size": 0, "total": 0, "total_pages": 1}, + "list": [], + "page": 1, + "page_size": 0, + "total": 0, } ) req = uburnode_somni_pb2.GetAudioReq(fetch_all=True) @@ -69,8 +85,10 @@ async def test_get_audio_passes_tag_code() -> None: rpc, service = _make_rpc() service.get_audio = AsyncMock( return_value={ - "materials": [], - "page": {"page": 1, "page_size": 20, "total": 0, "total_pages": 0}, + "list": [], + "page": 1, + "page_size": 20, + "total": 0, } ) req = uburnode_somni_pb2.GetAudioReq(tag_code="steady_rain") @@ -95,6 +113,9 @@ async def test_get_audio_tag_maps_fields() -> None: "code": "natural_sound", "name": "自然声", "name_en": "Natural Sound", + "id": "root-natural", + "parent_tag_id": None, + "status": "启用", } ] } @@ -103,11 +124,18 @@ async def test_get_audio_tag_maps_fields() -> None: service.get_audio_tag.assert_awaited_once_with() assert res.tags[0].code == "natural_sound" assert res.tags[0].name_en == "Natural Sound" + assert res.tags[0].id == "root-natural" + assert res.tags[0].parent_tag_id.WhichOneof("kind") == "null_value" + assert res.tags[0].status == "启用" @pytest.mark.asyncio -async def test_get_hot_no_args() -> None: +async def test_get_hot_maps_items() -> None: rpc, service = _make_rpc() + service.get_hot = AsyncMock( + return_value={"items": [{"keyword": "雨声", "score": 5}]} + ) res = await rpc.GetHot(uburnode_somni_pb2.GetHotReq(), _context()) service.get_hot.assert_awaited_once_with() - assert res == uburnode_somni_pb2.GetHotRes() + assert res.items[0].keyword == "雨声" + assert res.items[0].score == 5 diff --git a/tests/test_somni_audio_catalog.py b/tests/test_somni_audio_catalog.py index 7939c99..d47dcc9 100644 --- a/tests/test_somni_audio_catalog.py +++ b/tests/test_somni_audio_catalog.py @@ -2,6 +2,7 @@ from __future__ import annotations +import asyncio from unittest.mock import AsyncMock, MagicMock import pytest @@ -74,6 +75,7 @@ def _service( *, es_search: MagicMock | None = None, encoder: MagicMock | None = None, + hot: MagicMock | None = None, fetch_all_hard_limit: int = 50, cache_ttl_sec: float = 60.0, ) -> AudioCatalogService: @@ -87,7 +89,7 @@ def _service( default_page_size=1, max_page_size=200, fetch_all_hard_limit=fetch_all_hard_limit, - get_audio_root_tag_sim_threshold=0.85, + get_audio_root_tag_sim_threshold=0.75, somni_audio_catalog_cache_ttl_sec=cache_ttl_sec, ) return AudioCatalogService( @@ -95,6 +97,7 @@ def _service( settings, es_search=es_search, encoder=encoder, + hot=hot, ) @@ -115,9 +118,17 @@ async def test_get_audio_filters_content_form_code_then_pages() -> None: query_text="", tag_code="steady_rain", ) - assert [item["id"] for item in payload["materials"]] == ["a1"] - assert payload["page"]["total"] == 1 - assert "embedding" not in payload["materials"][0] + assert [item["id"] for item in payload["list"]] == ["a1"] + assert payload["total"] == 1 + assert set(payload["list"][0]) == { + "id", + "audio_name", + "audio_url", + "cover_url", + "description", + "vip", + } + assert payload["list"][0]["vip"] == 0 @pytest.mark.asyncio @@ -157,7 +168,7 @@ async def test_get_audio_keeps_mongo_and_es_caches_separate() -> None: ) es_search.list_audio_catalog_docs.assert_awaited_once_with(size=51) - assert [item["id"] for item in payload["materials"]] == ["a1"] + assert [item["id"] for item in payload["list"]] == ["a1"] @pytest.mark.asyncio @@ -211,7 +222,195 @@ async def test_get_audio_query_text_matches_root_tag_via_es() -> None: ) collection.find.assert_not_called() es_search.list_audio_catalog_docs.assert_awaited_once_with(size=51) - assert [item["id"] for item in payload["materials"]] == ["a1"] + assert [item["id"] for item in payload["list"]] == ["a1"] + + +@pytest.mark.asyncio +async def test_get_audio_query_text_matches_child_tag() -> None: + """搜索词更接近二级标签(如「白噪音」)时应按子标签命中,而非仅根标签。""" + white_noise = { + "_id": "a3", + "audio_name": "白噪音", + "content_form_tags": [ + { + "tag_id": "root-color-noise", + "code": "color_noise", + "name": "颜色噪音", + "parent_tag_id": None, + }, + { + "tag_id": "child-white-noise", + "code": "white_noise", + "name": "白噪音", + "parent_tag_id": "root-color-noise", + }, + ], + } + encoder = MagicMock() + encoder.is_loaded = True + encoder.encode_one = AsyncMock(return_value=[1.0, 0.0]) + es_search = MagicMock() + es_search.list_content_tag_vectors = AsyncMock( + return_value=[ + { + "id": "root-color-noise", + "dimension": "content_form", + "parent_tag_id": "", + "vector": [0.0, 1.0], # 与查询正交,根标签不命中 + }, + { + "id": "child-white-noise", + "dimension": "content_form", + "parent_tag_id": "root-color-noise", + "vector": [1.0, 0.0], # 子标签命中 + }, + ] + ) + es_search.list_audio_catalog_docs = AsyncMock(return_value=[_RAIN, white_noise]) + svc = _service( + _mongo_collection([white_noise]), + es_search=es_search, + encoder=encoder, + ) + payload = await svc.get_audio( + page=1, + page_size=10, + fetch_all=False, + query_text="白噪音", + tag_code="", + ) + assert [item["id"] for item in payload["list"]] == ["a3"] + + +@pytest.mark.asyncio +async def test_get_audio_query_text_returns_empty_when_child_unused() -> None: + """词典子标签命中但物料未挂该子标签时,直接返回空列表,不回退父级。""" + pink = { + "_id": "a4", + "audio_name": "粉噪音", + "content_form_tags": [ + { + "tag_id": "root-color-noise", + "code": "color_noise", + "name": "颜色噪音", + "parent_tag_id": None, + }, + { + "tag_id": "child-pink-noise", + "code": "pink_noise", + "name": "粉噪音", + "parent_tag_id": "root-color-noise", + }, + ], + } + encoder = MagicMock() + encoder.is_loaded = True + encoder.encode_one = AsyncMock(return_value=[1.0, 0.0]) + es_search = MagicMock() + es_search.list_content_tag_vectors = AsyncMock( + return_value=[ + { + "id": "root-color-noise", + "dimension": "content_form", + "parent_tag_id": "", + "vector": [0.2, 0.8], + }, + { + "id": "child-white-noise", + "dimension": "content_form", + "parent_tag_id": "root-color-noise", + "vector": [1.0, 0.0], + }, + ] + ) + es_search.list_audio_catalog_docs = AsyncMock(return_value=[_RAIN, pink]) + svc = _service(_mongo_collection([pink]), es_search=es_search, encoder=encoder) + payload = await svc.get_audio( + page=1, + page_size=10, + fetch_all=False, + query_text="白噪音", + tag_code="", + ) + assert payload["list"] == [] + assert payload["total"] == 0 + + +@pytest.mark.asyncio +async def test_get_audio_records_hot_when_query() -> None: + encoder = MagicMock() + encoder.is_loaded = True + encoder.encode_one = AsyncMock(return_value=[1.0, 0.0]) + es_search = MagicMock() + es_search.list_content_tag_vectors = AsyncMock( + return_value=[ + { + "id": "root-rain", + "dimension": "content_form", + "parent_tag_id": "", + "vector": [1.0, 0.0], + } + ] + ) + es_search.list_audio_catalog_docs = AsyncMock(return_value=[_RAIN]) + hot = MagicMock() + hot.record_search = AsyncMock() + svc = _service( + _mongo_collection([_RAIN]), + es_search=es_search, + encoder=encoder, + hot=hot, + ) + + await svc.get_audio( + page=1, + page_size=10, + fetch_all=False, + query_text=" 雨声 ", + tag_code="", + ) + await asyncio.sleep(0) + + hot.record_search.assert_awaited_once_with(" 雨声 ", hit_count=1) + + +@pytest.mark.asyncio +async def test_get_audio_succeeds_when_hot_recording_fails() -> None: + encoder = MagicMock() + encoder.is_loaded = True + encoder.encode_one = AsyncMock(return_value=[1.0, 0.0]) + es_search = MagicMock() + es_search.list_content_tag_vectors = AsyncMock( + return_value=[ + { + "id": "root-rain", + "dimension": "content_form", + "parent_tag_id": "", + "vector": [1.0, 0.0], + } + ] + ) + es_search.list_audio_catalog_docs = AsyncMock(return_value=[_RAIN]) + hot = MagicMock() + hot.record_search = AsyncMock(side_effect=RuntimeError("hot unavailable")) + svc = _service( + _mongo_collection([_RAIN]), + es_search=es_search, + encoder=encoder, + hot=hot, + ) + + payload = await svc.get_audio( + page=1, + page_size=10, + fetch_all=False, + query_text="雨声", + tag_code="", + ) + await asyncio.sleep(0) + + assert [item["id"] for item in payload["list"]] == ["a1"] + hot.record_search.assert_awaited_once() @pytest.mark.asyncio @@ -273,14 +472,94 @@ async def test_get_audio_tag_maps_root_fields() -> None: return_value=_Cursor( [ { + "_id": "root-natural", "type": "content_form", "code": "natural_sound", "name": "自然声", "name_en": "Natural Sound", + "parent_tag_id": None, + "status": "启用", } ] ) ) svc = _service(collection) payload = await svc.get_audio_tag() - assert payload["tags"][0]["code"] == "natural_sound" + assert payload["tags"][0] == { + "type": "content_form", + "code": "natural_sound", + "name": "自然声", + "name_en": "Natural Sound", + "id": "root-natural", + "parent_tag_id": None, + "status": "启用", + } + query = catalog._root_tag_query() + assert query["type"] == "content_form" + assert query["status"] == "启用" + collection.find.assert_called_once_with( + query, + { + "_id": 1, + "id": 1, + "type": 1, + "code": 1, + "name": 1, + "name_en": 1, + "parent_tag_id": 1, + "status": 1, + }, + ) + + +def test_root_tag_query_only_content_form() -> None: + query = catalog._root_tag_query() + assert query["type"] == "content_form" + + +def test_to_vip_normalizes_bool_int_and_string() -> None: + assert catalog._to_vip(None) == 0 + assert catalog._to_vip(False) == 0 + assert catalog._to_vip(0) == 0 + assert catalog._to_vip("false") == 0 + assert catalog._to_vip("0") == 0 + assert catalog._to_vip(True) == 1 + assert catalog._to_vip(1) == 1 + assert catalog._to_vip("true") == 1 + assert catalog._to_vip("1") == 1 + + +@pytest.mark.asyncio +async def test_get_audio_maps_vip_true_to_one() -> None: + doc = {**_RAIN, "vip": True} + svc = _service(_mongo_collection([doc])) + payload = await svc.get_audio( + page=1, + page_size=10, + fetch_all=False, + query_text="", + tag_code="", + ) + assert payload["list"][0]["vip"] == 1 + + +@pytest.mark.asyncio +async def test_drain_hot_tasks_waits_for_pending_recording() -> None: + started = asyncio.Event() + release = asyncio.Event() + + async def _slow_record(*_args, **_kwargs): + started.set() + await release.wait() + + hot = MagicMock() + hot.record_search = AsyncMock(side_effect=_slow_record) + svc = _service(_mongo_collection([_RAIN]), hot=hot) + svc._schedule_hot("雨声", 1) + + await started.wait() + release.set() + await svc.drain_hot_tasks(timeout_sec=1.0) + + assert not svc._hot_tasks + hot.record_search.assert_awaited_once_with("雨声", hit_count=1) diff --git a/tests/test_somni_audio_hot.py b/tests/test_somni_audio_hot.py new file mode 100644 index 0000000..0e6b0bd --- /dev/null +++ b/tests/test_somni_audio_hot.py @@ -0,0 +1,248 @@ +import asyncio +from unittest.mock import AsyncMock, MagicMock + +import pytest + +from app.core.config import Settings +from app.es.search_events import SearchEventsStore + + +def test_resolve_somni_redis_url_empty_does_not_fall_back() -> None: + from app.core.somni_redis import resolve_somni_redis_url + + settings = Settings(somni_redis_url=" ", redis_url="redis://localhost:6379/0") + assert resolve_somni_redis_url(settings) == "" + + +def test_resolve_somni_redis_url_prefers_somni() -> None: + from app.core.somni_redis import resolve_somni_redis_url + + settings = Settings( + somni_redis_url="redis://somni:6379/1", + redis_url="redis://localhost:6379/0", + ) + assert resolve_somni_redis_url(settings) == "redis://somni:6379/1" + + +@pytest.mark.asyncio +async def test_create_somni_redis_uses_only_somni_config(monkeypatch) -> None: + from app.core.somni_redis import create_somni_redis + + client = MagicMock() + client.ping = AsyncMock() + factory = MagicMock(return_value=client) + monkeypatch.setattr("app.core.somni_redis.Redis.from_url", factory) + settings = Settings( + somni_redis_url="redis://somni:6379/0", + redis_url="redis://shared:6379/0", + somni_redis_max_connections=64, + somni_redis_connect_timeout_sec=1.5, + somni_redis_socket_timeout_sec=2.5, + ) + + assert await create_somni_redis(settings) is client + + factory.assert_called_once_with( + "redis://somni:6379/0", + decode_responses=True, + max_connections=64, + socket_connect_timeout=1.5, + socket_timeout=2.5, + health_check_interval=30, + ) + client.ping.assert_awaited_once() + + +@pytest.mark.asyncio +async def test_create_somni_redis_closes_failed_client(monkeypatch) -> None: + from app.core.somni_redis import create_somni_redis + + client = MagicMock() + client.ping = AsyncMock(side_effect=ConnectionError("unavailable")) + client.aclose = AsyncMock() + monkeypatch.setattr("app.core.somni_redis.Redis.from_url", MagicMock(return_value=client)) + + with pytest.raises(ConnectionError): + await create_somni_redis(Settings(somni_redis_url="redis://somni:6379/0")) + + client.aclose.assert_awaited_once() + + +def test_hot_settings_defaults() -> None: + s = Settings() + assert s.somni_hot_enabled is True + assert s.somni_hot_top_n == 10 + assert s.somni_hot_redis_key == "somni:audio:hot:v1" + assert s.somni_redis_max_connections == 128 + assert s.somni_redis_connect_timeout_sec == 2.0 + assert s.somni_redis_socket_timeout_sec == 2.0 + assert s.somni_es_search_events_index == "somni_audio_search_events" + + +@pytest.mark.asyncio +async def test_search_events_ensure_and_index() -> None: + client = MagicMock() + client.indices.exists = AsyncMock(return_value=False) + client.indices.create = AsyncMock() + client.index = AsyncMock() + store = SearchEventsStore(client, Settings()) + await store.ensure_index() + client.indices.create.assert_awaited() + await store.index_event(keyword="雨声", raw_query=" 雨声 ", hit_count=2) + client.indices.exists.assert_awaited_once() + client.indices.create.assert_awaited_once() + kwargs = client.index.await_args.kwargs + assert kwargs["index"] == "somni_audio_search_events" + assert kwargs["document"]["keyword"] == "雨声" + assert kwargs["document"]["hit_count"] == 2 + + +@pytest.mark.asyncio +async def test_search_events_already_exists_still_indexes_and_caches_ensure() -> None: + class _AlreadyExistsError(Exception): + error = "resource_already_exists_exception" + + client = MagicMock() + client.indices.exists = AsyncMock(return_value=False) + client.indices.create = AsyncMock(side_effect=_AlreadyExistsError()) + client.index = AsyncMock() + store = SearchEventsStore(client, Settings()) + + await asyncio.gather( + store.index_event(keyword="雨声", raw_query=" 雨声 ", hit_count=1), + store.index_event(keyword="风声", raw_query="风声", hit_count=2), + ) + + client.indices.exists.assert_awaited_once() + client.indices.create.assert_awaited_once() + assert client.index.await_count == 2 + + +def test_normalize_keyword_strips() -> None: + from app.server.somni.audio.hot import normalize_keyword + + assert normalize_keyword(" 雨声 ") == "雨声" + assert normalize_keyword(" ") == "" + + +@pytest.mark.asyncio +async def test_record_and_list_hot() -> None: + from app.server.somni.audio.hot import HotTracker + + redis = MagicMock() + redis.zincrby = AsyncMock() + redis.zrevrange = AsyncMock( + return_value=[("雨声".encode(), 2.0), ("暴雨声".encode(), 1.0)] + ) + events = MagicMock() + events.index_event = AsyncMock() + tracker = HotTracker(redis, events, Settings()) + await tracker.record_search(" 雨声 ", hit_count=3) + redis.zincrby.assert_awaited_once() + events.index_event.assert_awaited_once() + items = await tracker.list_hot() + assert items == [{"keyword": "雨声", "score": 2}, {"keyword": "暴雨声", "score": 1}] + + +@pytest.mark.asyncio +async def test_record_blank_skipped() -> None: + from app.server.somni.audio.hot import HotTracker + + redis = MagicMock() + redis.zincrby = AsyncMock() + events = MagicMock() + events.index_event = AsyncMock() + tracker = HotTracker(redis, events, Settings()) + await tracker.record_search(" ", hit_count=0) + redis.zincrby.assert_not_called() + events.index_event.assert_not_called() + + +@pytest.mark.asyncio +async def test_record_search_redis_failure_still_indexes_es() -> None: + from app.server.somni.audio.hot import HotTracker + + redis = MagicMock() + redis.zincrby = AsyncMock(side_effect=RuntimeError("redis unavailable")) + events = MagicMock() + events.index_event = AsyncMock() + tracker = HotTracker(redis, events, Settings()) + + await tracker.record_search(" 雨声 ", hit_count=3) + + events.index_event.assert_awaited_once_with( + keyword="雨声", + raw_query=" 雨声 ", + hit_count=3, + ) + + +@pytest.mark.asyncio +async def test_record_search_es_failure_returns_after_redis_increment() -> None: + from app.server.somni.audio.hot import HotTracker + + redis = MagicMock() + redis.zincrby = AsyncMock() + events = MagicMock() + events.index_event = AsyncMock(side_effect=RuntimeError("es unavailable")) + tracker = HotTracker(redis, events, Settings()) + + await tracker.record_search("雨声", hit_count=1) + + redis.zincrby.assert_awaited_once() + + +@pytest.mark.asyncio +async def test_hot_disabled_skips_writes_and_returns_empty_list() -> None: + from app.server.somni.audio.hot import HotTracker + + redis = MagicMock() + redis.zincrby = AsyncMock() + redis.zrevrange = AsyncMock() + events = MagicMock() + events.index_event = AsyncMock() + tracker = HotTracker(redis, events, Settings(somni_hot_enabled=False)) + + await tracker.record_search("雨声", hit_count=1) + assert await tracker.list_hot() == [] + + redis.zincrby.assert_not_called() + redis.zrevrange.assert_not_called() + events.index_event.assert_not_called() + + +@pytest.mark.asyncio +async def test_list_hot_non_positive_top_n_returns_empty_without_redis_read() -> None: + from app.server.somni.audio.hot import HotTracker + + redis = MagicMock() + redis.zrevrange = AsyncMock() + tracker = HotTracker(redis, None, Settings(somni_hot_top_n=0)) + + assert await tracker.list_hot() == [] + redis.zrevrange.assert_not_called() + + +@pytest.mark.asyncio +async def test_list_hot_requires_redis() -> None: + from app.core.exceptions import AppError + from app.server.somni.audio.hot import HotTracker + + tracker = HotTracker(None, None, Settings()) + with pytest.raises(AppError): + await tracker.list_hot() + + +@pytest.mark.asyncio +async def test_list_hot_redis_failure_is_service_unavailable() -> None: + from app.core.exceptions import AppError + from app.server.somni.audio.hot import HotTracker + + redis = MagicMock() + redis.zrevrange = AsyncMock(side_effect=ConnectionError("unavailable")) + tracker = HotTracker(redis, None, Settings()) + + with pytest.raises(AppError) as exc: + await tracker.list_hot() + + assert exc.value.status_code == 503 diff --git a/tests/test_uburnode_proto_import.py b/tests/test_uburnode_proto_import.py index 69ce72b..285108b 100644 --- a/tests/test_uburnode_proto_import.py +++ b/tests/test_uburnode_proto_import.py @@ -1,5 +1,7 @@ """proto gen 产物可 import。""" +from google.protobuf.struct_pb2 import Value + from app.uburnode_grpc.grpc_gen import ( uburnode_pb2, uburnode_pb2_grpc, @@ -25,3 +27,46 @@ def test_somni_package() -> None: answers_field = uburnode_somni_pb2.GetAnswerRes.DESCRIPTOR.fields_by_name["answers"] assert answers_field.type == answers_field.TYPE_STRING assert not answers_field.is_repeated + + +def test_get_hot_res_has_items() -> None: + from app.uburnode_grpc.grpc_gen import uburnode_somni_pb2 + + res = uburnode_somni_pb2.GetHotRes( + items=[uburnode_somni_pb2.HotKeyword(keyword="雨声", score=3)] + ) + assert res.items[0].keyword == "雨声" + assert res.items[0].score == 3 + + +def test_tag_dict_item_has_id_parent_status() -> None: + from app.uburnode_grpc.grpc_gen import uburnode_somni_pb2 + + item = uburnode_somni_pb2.TagDictItem( + id="t1", + parent_tag_id=Value(null_value=0), + status="启用", + type="content_form", + code="rain", + name="雨声", + name_en="Rain", + ) + assert item.id == "t1" + assert item.parent_tag_id.WhichOneof("kind") == "null_value" + assert item.status == "启用" + + +def test_audio_list_item_keeps_default_field_presence() -> None: + item = uburnode_somni_pb2.AudioListItem( + id="m1", + audio_name="雨声", + audio_url="", + cover_url="", + description="", + vip=0, + ) + assert item.HasField("audio_url") + assert item.HasField("cover_url") + assert item.HasField("description") + assert item.HasField("vip") + assert item.vip == 0