213 lines
7.0 KiB
Python
213 lines
7.0 KiB
Python
"""Pipeline configuration and statistics service for VoIdea."""
|
|
|
|
import json
|
|
from collections import Counter
|
|
from datetime import datetime, timezone
|
|
from typing import Any, Optional
|
|
|
|
from sqlalchemy import func, select
|
|
from sqlalchemy.ext.asyncio import AsyncSession
|
|
|
|
from app.models.agent import AgentConfig
|
|
from app.models.pipeline import PipelineStats
|
|
from app.schemas.pipeline import PipelineConfigResponse, PipelineStageConfig
|
|
|
|
DEFAULT_PIPELINE_CONFIG: dict[str, Any] = {
|
|
"vad": {
|
|
"enabled": True,
|
|
"noise_threshold": 0.3,
|
|
"silence_timeout_ms": 1500,
|
|
"min_audio_duration_ms": 300,
|
|
},
|
|
"wake_word": {
|
|
"enabled": True,
|
|
"word": "ВоИдея",
|
|
"timeout_minutes": 5,
|
|
"sensitivity": 0.7,
|
|
},
|
|
"semantic_validation": {
|
|
"mode": "fast",
|
|
"timeout_ms": 5000,
|
|
},
|
|
"confidence": {
|
|
"verified_threshold": 80,
|
|
"warning_threshold": 50,
|
|
},
|
|
"agent_chaining": {
|
|
"max_agents_per_dialog": 3,
|
|
"enabled": True,
|
|
},
|
|
"auto_tuning": {
|
|
"enabled": True,
|
|
"min_samples": 5,
|
|
"learning_rate": 0.1,
|
|
},
|
|
"stages_order": [
|
|
"vad",
|
|
"wake_word",
|
|
"semantic_validation",
|
|
"routing",
|
|
"verification",
|
|
],
|
|
}
|
|
|
|
|
|
def _dict_to_response(data: dict[str, Any]) -> PipelineConfigResponse:
|
|
stages = data.get("stages_order", DEFAULT_PIPELINE_CONFIG["stages_order"])
|
|
config_data = {k: v for k, v in data.items() if k != "stages_order"}
|
|
|
|
def stage_or_default(name: str) -> PipelineStageConfig:
|
|
default = DEFAULT_PIPELINE_CONFIG.get(name, {})
|
|
override = config_data.get(name, {})
|
|
merged = {**default, **override}
|
|
return PipelineStageConfig(**merged)
|
|
|
|
return PipelineConfigResponse(
|
|
vad=stage_or_default("vad"),
|
|
wake_word=stage_or_default("wake_word"),
|
|
semantic_validation=stage_or_default("semantic_validation"),
|
|
confidence=stage_or_default("confidence"),
|
|
agent_chaining=stage_or_default("agent_chaining"),
|
|
auto_tuning=stage_or_default("auto_tuning"),
|
|
stages_order=stages,
|
|
)
|
|
|
|
|
|
class PipelineService:
|
|
def __init__(self, db: AsyncSession):
|
|
self.db = db
|
|
|
|
async def get_config(self) -> PipelineConfigResponse:
|
|
result = await self.db.execute(
|
|
select(AgentConfig).where(AgentConfig.agent_name == "conductor")
|
|
)
|
|
conductor = result.scalar_one_or_none()
|
|
if not conductor or not conductor.config:
|
|
return _dict_to_response(dict(DEFAULT_PIPELINE_CONFIG))
|
|
|
|
try:
|
|
data = json.loads(conductor.config)
|
|
if not isinstance(data, dict):
|
|
return _dict_to_response(dict(DEFAULT_PIPELINE_CONFIG))
|
|
return _dict_to_response(data)
|
|
except (json.JSONDecodeError, TypeError):
|
|
return _dict_to_response(dict(DEFAULT_PIPELINE_CONFIG))
|
|
|
|
async def update_config(self, updates: dict[str, Any]) -> PipelineConfigResponse:
|
|
result = await self.db.execute(
|
|
select(AgentConfig).where(AgentConfig.agent_name == "conductor")
|
|
)
|
|
conductor = result.scalar_one_or_none()
|
|
if not conductor:
|
|
conductor = AgentConfig(
|
|
agent_name="conductor",
|
|
description="Дирижёр — главный оркестратор",
|
|
is_enabled=True,
|
|
version="1.0.0",
|
|
config=json.dumps(DEFAULT_PIPELINE_CONFIG, ensure_ascii=False),
|
|
)
|
|
self.db.add(conductor)
|
|
|
|
current = {}
|
|
if conductor.config:
|
|
try:
|
|
current = json.loads(conductor.config)
|
|
except (json.JSONDecodeError, TypeError):
|
|
current = {}
|
|
|
|
merged = {**DEFAULT_PIPELINE_CONFIG, **current}
|
|
|
|
stages_order = updates.pop("stages_order", None)
|
|
if stages_order is not None:
|
|
merged["stages_order"] = stages_order
|
|
|
|
for key, value in updates.items():
|
|
if value is not None and isinstance(value, dict):
|
|
existing = merged.get(key, {})
|
|
if isinstance(existing, dict):
|
|
merged[key] = {**existing, **value}
|
|
else:
|
|
merged[key] = value
|
|
|
|
conductor.config = json.dumps(merged, ensure_ascii=False)
|
|
await self.db.commit()
|
|
|
|
return _dict_to_response(merged)
|
|
|
|
async def record_stat(
|
|
self,
|
|
user_id: str | None,
|
|
stage: str,
|
|
passed: bool,
|
|
reason: str | None = None,
|
|
duration_ms: int | None = None,
|
|
) -> PipelineStats:
|
|
stat = PipelineStats(
|
|
id=str(__import__("uuid").uuid4()),
|
|
user_id=user_id,
|
|
stage=stage,
|
|
passed=passed,
|
|
reason=reason,
|
|
duration_ms=duration_ms,
|
|
created_at=datetime.now(timezone.utc),
|
|
)
|
|
self.db.add(stat)
|
|
await self.db.commit()
|
|
return stat
|
|
|
|
async def list_stats(
|
|
self,
|
|
user_id: str | None = None,
|
|
stage: str | None = None,
|
|
skip: int = 0,
|
|
limit: int = 50,
|
|
) -> list[PipelineStats]:
|
|
query = select(PipelineStats).order_by(PipelineStats.created_at.desc())
|
|
if user_id:
|
|
query = query.where(PipelineStats.user_id == user_id)
|
|
if stage:
|
|
query = query.where(PipelineStats.stage == stage)
|
|
result = await self.db.execute(query.offset(skip).limit(limit))
|
|
return list(result.scalars().all())
|
|
|
|
async def get_summary(self) -> dict[str, Any]:
|
|
total = await self.db.execute(select(func.count(PipelineStats.id)))
|
|
total_count = total.scalar() or 0
|
|
|
|
passed_count = await self.db.execute(
|
|
select(func.count(PipelineStats.id)).where(PipelineStats.passed == True)
|
|
)
|
|
passed_total = passed_count.scalar() or 0
|
|
|
|
stage_result = await self.db.execute(
|
|
select(PipelineStats.stage, func.count(PipelineStats.id))
|
|
.group_by(PipelineStats.stage)
|
|
.order_by(func.count(PipelineStats.id).desc())
|
|
)
|
|
stages = {row[0]: row[1] for row in stage_result.all()}
|
|
|
|
avg_dur = await self.db.execute(
|
|
select(func.avg(PipelineStats.duration_ms)).where(PipelineStats.duration_ms.isnot(None))
|
|
)
|
|
avg_val = avg_dur.scalar()
|
|
|
|
fail_reason_result = await self.db.execute(
|
|
select(PipelineStats.reason, func.count(PipelineStats.id))
|
|
.where(
|
|
PipelineStats.passed == False,
|
|
PipelineStats.reason.isnot(None),
|
|
)
|
|
.group_by(PipelineStats.reason)
|
|
.order_by(func.count(PipelineStats.id).desc())
|
|
.limit(10)
|
|
)
|
|
fail_reasons = [(row[0], row[1]) for row in fail_reason_result.all()]
|
|
|
|
return {
|
|
"total_entries": total_count,
|
|
"stages": stages,
|
|
"passed_ratio": round(passed_total / total_count, 4) if total_count > 0 else 0.0,
|
|
"avg_duration_ms": round(float(avg_val), 1) if avg_val else None,
|
|
"fail_reasons": fail_reasons,
|
|
}
|