From cae8e6413764c50a2cd8a9bb3821a7c24fc4d3d0 Mon Sep 17 00:00:00 2001 From: huaiyj <8699003+huaiyj@user.noreply.gitee.com> Date: Mon, 24 Aug 2026 20:53:08 +0800 Subject: [PATCH] =?UTF-8?q?feat:=20=E6=8C=81=E4=B9=85=E5=8C=96=20LLM=20?= =?UTF-8?q?=E7=94=A8=E9=87=8F=E6=8C=87=E6=A0=87?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- ocai-service/ocai/ai/metrics.py | 41 +++++++- ocai-service/ocai/api/deps.py | 10 +- ocai-service/ocai/api/v2_route.py | 31 +++++- ocai-service/ocai/main.py | 16 +++- ocai-service/ocai/repo/mysql/__init__.py | 2 + .../migrations/versions/0002_usage_records.py | 42 ++++++++ ocai-service/ocai/repo/mysql/models.py | 23 +++++ ocai-service/ocai/repo/mysql/usage_repo.py | 96 +++++++++++++++++++ ocai-service/ocai/schemas/config.py | 10 ++ ocai-service/ocai/services/__init__.py | 2 + ocai-service/ocai/services/chat_service.py | 17 +++- .../ocai/services/usage_metrics_service.py | 39 ++++++++ ocai-service/sql/0002_usage_records.down.sql | 1 + ocai-service/sql/0002_usage_records.up.sql | 16 ++++ .../tests/ai/test_metrics_persistence.py | 36 +++++++ .../tests/api/test_usage_metrics_endpoint.py | 34 +++++++ ocai-service/tests/repo/test_mysql_repos.py | 47 +++++++++ 17 files changed, 456 insertions(+), 7 deletions(-) create mode 100644 ocai-service/ocai/repo/mysql/migrations/versions/0002_usage_records.py create mode 100644 ocai-service/ocai/repo/mysql/usage_repo.py create mode 100644 ocai-service/ocai/services/usage_metrics_service.py create mode 100644 ocai-service/sql/0002_usage_records.down.sql create mode 100644 ocai-service/sql/0002_usage_records.up.sql create mode 100644 ocai-service/tests/ai/test_metrics_persistence.py create mode 100644 ocai-service/tests/api/test_usage_metrics_endpoint.py diff --git a/ocai-service/ocai/ai/metrics.py b/ocai-service/ocai/ai/metrics.py index c5a5917..c02a1dc 100644 --- a/ocai-service/ocai/ai/metrics.py +++ b/ocai-service/ocai/ai/metrics.py @@ -19,6 +19,7 @@ import logging import threading import time from collections import deque +from concurrent.futures import ThreadPoolExecutor from dataclasses import dataclass, field from typing import Any @@ -39,6 +40,8 @@ class RequestMetrics: session_id: str = "" user_id: str = "" request_type: str = "" # "v2_chat" / "v2_chat_direct" + project_id: str = "" + model: str = "" # 时间节点(epoch seconds) request_received_at: float = field(default_factory=time.time) @@ -58,6 +61,7 @@ class RequestMetrics: # Token usage input_tokens: int = 0 output_tokens: int = 0 + estimated_cost: float = 0.0 # Error has_error: bool = False @@ -65,6 +69,7 @@ class RequestMetrics: # ReAct react_steps: int = 0 + _recorded: bool = field(default=False, repr=False) # ---- 记录方法 ---- @@ -144,7 +149,12 @@ class MetricsCollector: 线程安全,可在多 worker 进程中各自独立运行(进程内聚合)。 """ - def __init__(self, window_size: int = _WINDOW_SIZE): + def __init__( + self, + window_size: int = _WINDOW_SIZE, + persist_fn=None, + model_pricing: dict[str, tuple[float, float]] | None = None, + ): self._lock = threading.Lock() self._window: deque[dict[str, Any]] = deque(maxlen=window_size) @@ -152,14 +162,30 @@ class MetricsCollector: self._total_requests: int = 0 self._total_errors: int = 0 self._started_at: float = time.time() + self._persist_fn = persist_fn + self._model_pricing = model_pricing or {} + self._persist_executor = ( + ThreadPoolExecutor(max_workers=1, thread_name_prefix="usage-metrics") if persist_fn else None + ) def record(self, metrics: RequestMetrics) -> None: """将一次请求的指标写入聚合器""" + if metrics._recorded: + return + metrics._recorded = True metrics.finalize() + if metrics.estimated_cost == 0 and metrics.model in self._model_pricing: + input_price, output_price = self._model_pricing[metrics.model] + metrics.estimated_cost = ( + metrics.input_tokens * input_price + metrics.output_tokens * output_price + ) / 1_000_000 snapshot = metrics.to_dict() snapshot["session_id"] = metrics.session_id snapshot["user_id"] = metrics.user_id snapshot["request_type"] = metrics.request_type + snapshot["project_id"] = metrics.project_id + snapshot["model"] = metrics.model + snapshot["estimated_cost"] = metrics.estimated_cost snapshot["timestamp"] = metrics.ended_at or time.time() with self._lock: @@ -168,6 +194,19 @@ class MetricsCollector: if metrics.has_error: self._total_errors += 1 + if self._persist_executor is not None: + self._persist_executor.submit(self._persist_snapshot, dict(snapshot)) + + def _persist_snapshot(self, snapshot: dict[str, Any]) -> None: + try: + self._persist_fn(snapshot) + except Exception: + logger.exception("Failed to persist usage metrics") + + def shutdown(self) -> None: + if self._persist_executor is not None: + self._persist_executor.shutdown(wait=True) + def get_summary(self, last_n: int | None = None) -> dict[str, Any]: """返回聚合指标摘要 diff --git a/ocai-service/ocai/api/deps.py b/ocai-service/ocai/api/deps.py index e977074..0d43e9f 100644 --- a/ocai-service/ocai/api/deps.py +++ b/ocai-service/ocai/api/deps.py @@ -27,7 +27,7 @@ from fastapi import HTTPException, Request if TYPE_CHECKING: from ocai.schemas.config import ApplicationConfig - from ocai.services import ChatService, UserPreferencesService + from ocai.services import ChatService, UsageMetricsService, UserPreferencesService def get_chat_service(request: Request) -> ChatService: @@ -50,6 +50,13 @@ def get_user_preferences_service(request: Request) -> UserPreferencesService: return svc # type: ignore[no-any-return] +def get_usage_metrics_service(request: Request) -> UsageMetricsService: + svc = getattr(request.app.state, "usage_metrics_service", None) + if svc is None: + raise HTTPException(status_code=500, detail="UsageMetricsService not initialized") + return svc # type: ignore[no-any-return] + + def get_app_config(request: Request) -> ApplicationConfig | None: """Resolve the loaded :class:`ApplicationConfig` from ``app.state``. @@ -63,5 +70,6 @@ def get_app_config(request: Request) -> ApplicationConfig | None: __all__ = [ "get_app_config", "get_chat_service", + "get_usage_metrics_service", "get_user_preferences_service", ] diff --git a/ocai-service/ocai/api/v2_route.py b/ocai-service/ocai/api/v2_route.py index 6280b17..2fd80a3 100644 --- a/ocai-service/ocai/api/v2_route.py +++ b/ocai-service/ocai/api/v2_route.py @@ -32,13 +32,19 @@ from __future__ import annotations import logging from collections.abc import AsyncGenerator, Iterable +from datetime import datetime from typing import TYPE_CHECKING, Annotated, Any from fastapi import APIRouter, Body, Depends, Query, Request from fastapi.responses import JSONResponse, StreamingResponse from starlette.concurrency import iterate_in_threadpool -from ocai.api.deps import get_app_config, get_chat_service, get_user_preferences_service +from ocai.api.deps import ( + get_app_config, + get_chat_service, + get_usage_metrics_service, + get_user_preferences_service, +) from ocai.permission import UserInfo, get_current_user from ocai.schemas.agent_sse import AgentChatRequest from ocai.schemas.v2_ops import ( @@ -50,7 +56,7 @@ from ocai.schemas.v2_ops import ( if TYPE_CHECKING: from ocai.schemas.config import ApplicationConfig - from ocai.services import ChatService, UserPreferencesService + from ocai.services import ChatService, UsageMetricsService, UserPreferencesService logger = logging.getLogger(__name__) @@ -64,6 +70,7 @@ router = APIRouter(prefix="/api/v2/ocai", tags=["v2"]) ChatSvc = Annotated["ChatService", Depends(get_chat_service)] PrefSvc = Annotated["UserPreferencesService", Depends(get_user_preferences_service)] AppConf = Annotated["ApplicationConfig | None", Depends(get_app_config)] +UsageSvc = Annotated["UsageMetricsService", Depends(get_usage_metrics_service)] CurrentUser = Annotated[UserInfo, Depends(get_current_user)] @@ -443,3 +450,23 @@ async def get_recent_metrics( ) -> dict[str, Any]: recent = chat_service.metrics_collector.get_recent(n=n, user_id=user.login_name) return {"requests": recent, "count": len(recent)} + + +@router.get("/metrics/usage") +async def get_usage_metrics( + usage_service: UsageSvc, + user: CurrentUser, + project_id: str | None = Query(default=None), + model: str | None = Query(default=None), + start: datetime | None = Query(default=None), + end: datetime | None = Query(default=None), + group_by: str = Query(default="model", pattern="^(model|project|user|day|week|month)$"), +) -> dict[str, Any]: + return usage_service.query( + user_id=user.login_name, + project_id=project_id, + model=model, + start=start, + end=end, + group_by=group_by, + ) diff --git a/ocai-service/ocai/main.py b/ocai-service/ocai/main.py index 09e8a99..7d07ba8 100644 --- a/ocai-service/ocai/main.py +++ b/ocai-service/ocai/main.py @@ -25,6 +25,7 @@ from ocai.permission import PermissionConfig, setup_permission from ocai.repo.cache_repo import CacheRepository from ocai.repo.mysql import ( MysqlAgentRepository, + MysqlUsageRepository, MysqlUserPreferencesRepo, MysqlV2SessionRepository, ) @@ -32,6 +33,7 @@ from ocai.schemas.config import ApplicationConfig from ocai.services import ( ChatService, MessageConverter, + UsageMetricsService, UserPreferencesService, ) from ocai.utils.config_loader import config_source_label, load_application_from_env @@ -200,6 +202,7 @@ def create_app(conf: ApplicationConfig | None = None) -> FastAPI: v2_session_repo = MysqlV2SessionRepository(engine) user_preferences_repo = MysqlUserPreferencesRepo(engine) + usage_repo = MysqlUsageRepository(engine) user_preferences_repo.ensure_indexes() # Upstream auth provider: OCManagerClient when dependence.oc_manager.host @@ -219,7 +222,15 @@ def create_app(conf: ApplicationConfig | None = None) -> FastAPI: ) message_converter = MessageConverter() - metrics_collector = MetricsCollector() + model_pricing = {} + for llm in conf.dependence.llms: + prices = (llm.input_price_per_million, llm.output_price_per_million) + model_pricing[llm.name] = prices + model_pricing[llm.model] = prices + metrics_collector = MetricsCollector( + persist_fn=usage_repo.save_request_metrics, + model_pricing=model_pricing, + ) chat_service = ChatService( agent_manager=agent_manager, v2_session_repo=v2_session_repo, @@ -231,6 +242,7 @@ def create_app(conf: ApplicationConfig | None = None) -> FastAPI: user_preferences_service = UserPreferencesService( user_preferences_repo=user_preferences_repo, ) + usage_metrics_service = UsageMetricsService(usage_repo) @asynccontextmanager async def lifespan(_: FastAPI) -> AsyncIterator[None]: @@ -247,6 +259,7 @@ def create_app(conf: ApplicationConfig | None = None) -> FastAPI: agent_manager.shutdown() except Exception: logger.exception("AgentManager shutdown failed") + metrics_collector.shutdown() app = FastAPI( title="ocai-service", @@ -266,6 +279,7 @@ def create_app(conf: ApplicationConfig | None = None) -> FastAPI: app.state.metrics_collector = metrics_collector app.state.chat_service = chat_service app.state.user_preferences_service = user_preferences_service + app.state.usage_metrics_service = usage_metrics_service register_exception_handlers(app, logger=logger) diff --git a/ocai-service/ocai/repo/mysql/__init__.py b/ocai-service/ocai/repo/mysql/__init__.py index 4ae19b2..d3e56fb 100644 --- a/ocai-service/ocai/repo/mysql/__init__.py +++ b/ocai-service/ocai/repo/mysql/__init__.py @@ -8,6 +8,7 @@ implementations. from ocai.repo.mysql.agent_repo import MysqlAgentRepository from ocai.repo.mysql.chat_session_repo import MysqlChatSessionRepository from ocai.repo.mysql.engine import MySQLEngine +from ocai.repo.mysql.usage_repo import MysqlUsageRepository from ocai.repo.mysql.user_preferences_repo import MysqlUserPreferencesRepo from ocai.repo.mysql.v2_session_repo import MysqlV2SessionRepository @@ -15,6 +16,7 @@ __all__ = [ "MySQLEngine", "MysqlAgentRepository", "MysqlChatSessionRepository", + "MysqlUsageRepository", "MysqlUserPreferencesRepo", "MysqlV2SessionRepository", ] diff --git a/ocai-service/ocai/repo/mysql/migrations/versions/0002_usage_records.py b/ocai-service/ocai/repo/mysql/migrations/versions/0002_usage_records.py new file mode 100644 index 0000000..0e8af64 --- /dev/null +++ b/ocai-service/ocai/repo/mysql/migrations/versions/0002_usage_records.py @@ -0,0 +1,42 @@ +"""persistent LLM usage records + +Revision ID: 0002_usage +Revises: 0001_initial +""" + +from collections.abc import Sequence + +import sqlalchemy as sa +from alembic import op + +revision: str = "0002_usage" +down_revision: str | None = "0001_initial" +branch_labels: str | Sequence[str] | None = None +depends_on: str | Sequence[str] | None = None + + +def upgrade() -> None: + op.create_table( + "usage_records", + sa.Column("id", sa.String(length=64), nullable=False), + sa.Column("user_id", sa.String(length=128), nullable=False), + sa.Column("session_id", sa.String(length=64), nullable=False), + sa.Column("project_id", sa.String(length=128), nullable=True), + sa.Column("model", sa.String(length=128), nullable=False), + sa.Column("input_tokens", sa.Integer(), nullable=False), + sa.Column("output_tokens", sa.Integer(), nullable=False), + sa.Column("duration_ms", sa.Float(), nullable=False), + sa.Column("estimated_cost", sa.Float(), nullable=False), + sa.Column("created_at", sa.DateTime(), nullable=False), + sa.PrimaryKeyConstraint("id"), + ) + op.create_index("idx_usage_user_created", "usage_records", ["user_id", "created_at"]) + op.create_index("idx_usage_project_created", "usage_records", ["project_id", "created_at"]) + op.create_index("idx_usage_model_created", "usage_records", ["model", "created_at"]) + + +def downgrade() -> None: + op.drop_index("idx_usage_model_created", table_name="usage_records") + op.drop_index("idx_usage_project_created", table_name="usage_records") + op.drop_index("idx_usage_user_created", table_name="usage_records") + op.drop_table("usage_records") diff --git a/ocai-service/ocai/repo/mysql/models.py b/ocai-service/ocai/repo/mysql/models.py index 415e185..adad28f 100644 --- a/ocai-service/ocai/repo/mysql/models.py +++ b/ocai-service/ocai/repo/mysql/models.py @@ -29,6 +29,7 @@ from sqlalchemy import ( JSON, BigInteger, DateTime, + Float, Index, Integer, String, @@ -195,11 +196,33 @@ class UserPreference(Base): updated_at: Mapped[datetime] = mapped_column(DateTime, nullable=False) +class UsageRecord(Base): + __tablename__ = "usage_records" + + id: Mapped[str] = mapped_column(String(64), primary_key=True) + user_id: Mapped[str] = mapped_column(String(128), nullable=False) + session_id: Mapped[str] = mapped_column(String(64), nullable=False) + project_id: Mapped[str | None] = mapped_column(String(128)) + model: Mapped[str] = mapped_column(String(128), nullable=False, default="") + input_tokens: Mapped[int] = mapped_column(Integer, nullable=False, default=0) + output_tokens: Mapped[int] = mapped_column(Integer, nullable=False, default=0) + duration_ms: Mapped[float] = mapped_column(Float, nullable=False, default=0) + estimated_cost: Mapped[float] = mapped_column(Float, nullable=False, default=0) + created_at: Mapped[datetime] = mapped_column(DateTime, nullable=False) + + __table_args__ = ( + Index("idx_usage_user_created", "user_id", "created_at"), + Index("idx_usage_project_created", "project_id", "created_at"), + Index("idx_usage_model_created", "model", "created_at"), + ) + + __all__ = [ "AgentLog", "AgentMemory", "Base", "ChatSession", + "UsageRecord", "UserPreference", "V2Session", "V2SessionInvolvedAgent", diff --git a/ocai-service/ocai/repo/mysql/usage_repo.py b/ocai-service/ocai/repo/mysql/usage_repo.py new file mode 100644 index 0000000..ac419b8 --- /dev/null +++ b/ocai-service/ocai/repo/mysql/usage_repo.py @@ -0,0 +1,96 @@ +from __future__ import annotations + +import uuid +from datetime import UTC, datetime +from typing import Any + +from sqlalchemy import func, select + +from ocai.repo.mysql.engine import MySQLEngine +from ocai.repo.mysql.models import UsageRecord + + +class MysqlUsageRepository: + def __init__(self, engine: MySQLEngine): + self._engine = engine + + def save_request_metrics(self, snapshot: dict[str, Any]) -> None: + created_at = datetime.fromtimestamp( + float(snapshot.get("timestamp") or datetime.now(UTC).timestamp()), + tz=UTC, + ).replace(tzinfo=None) + row = UsageRecord( + id=uuid.uuid4().hex, + user_id=str(snapshot.get("user_id") or ""), + session_id=str(snapshot.get("session_id") or ""), + project_id=str(snapshot.get("project_id") or "") or None, + model=str(snapshot.get("model") or ""), + input_tokens=max(0, int(snapshot.get("input_tokens") or 0)), + output_tokens=max(0, int(snapshot.get("output_tokens") or 0)), + duration_ms=max(0.0, float(snapshot.get("total_duration_ms") or 0)), + estimated_cost=max(0.0, float(snapshot.get("estimated_cost") or 0)), + created_at=created_at, + ) + with self._engine.session_scope() as session: + session.add(row) + + def aggregate( + self, + *, + user_id: str, + project_id: str | None = None, + model: str | None = None, + start: datetime | None = None, + end: datetime | None = None, + group_by: str = "model", + ) -> list[dict[str, Any]]: + dialect = self._engine.engine.dialect.name + if dialect == "mysql": + week_group = func.date_format(UsageRecord.created_at, "%x-W%v") + month_group = func.date_format(UsageRecord.created_at, "%Y-%m") + else: + week_group = func.strftime("%Y-W%W", UsageRecord.created_at) + month_group = func.strftime("%Y-%m", UsageRecord.created_at) + group_column = { + "model": UsageRecord.model, + "project": UsageRecord.project_id, + "user": UsageRecord.user_id, + "day": func.date(UsageRecord.created_at), + "week": week_group, + "month": month_group, + }[group_by] + conditions = [UsageRecord.user_id == user_id] + if project_id: + conditions.append(UsageRecord.project_id == project_id) + if model: + conditions.append(UsageRecord.model == model) + if start: + conditions.append(UsageRecord.created_at >= start) + if end: + conditions.append(UsageRecord.created_at <= end) + stmt = ( + select( + group_column.label("group"), + func.count(UsageRecord.id).label("requests"), + func.sum(UsageRecord.input_tokens).label("input_tokens"), + func.sum(UsageRecord.output_tokens).label("output_tokens"), + func.sum(UsageRecord.estimated_cost).label("estimated_cost"), + func.avg(UsageRecord.duration_ms).label("avg_duration_ms"), + ) + .where(*conditions) + .group_by(group_column) + .order_by(group_column) + ) + with self._engine.session_scope() as session: + rows = session.execute(stmt).all() + return [ + { + "group": str(row.group or ""), + "requests": int(row.requests or 0), + "input_tokens": int(row.input_tokens or 0), + "output_tokens": int(row.output_tokens or 0), + "estimated_cost": round(float(row.estimated_cost or 0), 6), + "avg_duration_ms": round(float(row.avg_duration_ms or 0), 1), + } + for row in rows + ] diff --git a/ocai-service/ocai/schemas/config.py b/ocai-service/ocai/schemas/config.py index 1dd360a..ce64b22 100644 --- a/ocai-service/ocai/schemas/config.py +++ b/ocai-service/ocai/schemas/config.py @@ -62,6 +62,16 @@ class LLMConfig(BaseModel): default=False, description="前端 deep-think 开关默认值;仅作 UI 提示,不影响后端调用", ) + input_price_per_million: float = Field( + default=0, + ge=0, + description="每百万 input tokens 的估算价格;未配置时成本记为 0", + ) + output_price_per_million: float = Field( + default=0, + ge=0, + description="每百万 output tokens 的估算价格;未配置时成本记为 0", + ) default: bool = Field( default=False, description=( diff --git a/ocai-service/ocai/services/__init__.py b/ocai-service/ocai/services/__init__.py index 98c8e1b..55e736d 100644 --- a/ocai-service/ocai/services/__init__.py +++ b/ocai-service/ocai/services/__init__.py @@ -1,9 +1,11 @@ from .chat_service import ChatService from .message_converter import MessageConverter +from .usage_metrics_service import UsageMetricsService from .user_preferences_service import UserPreferencesService __all__ = [ "ChatService", "MessageConverter", + "UsageMetricsService", "UserPreferencesService", ] diff --git a/ocai-service/ocai/services/chat_service.py b/ocai-service/ocai/services/chat_service.py index 0d9d45f..79e7c11 100644 --- a/ocai-service/ocai/services/chat_service.py +++ b/ocai-service/ocai/services/chat_service.py @@ -89,6 +89,15 @@ class ChatService(MessagePersistenceMixin): def metrics_collector(self) -> MetricsCollector: return self._metrics + def _metrics_model_name(self, requested_model: str) -> str: + if requested_model: + return requested_model + try: + model_name = self._agent_manager.get_llm_client().model_name + return model_name if isinstance(model_name, str) else "" + except Exception: + return "" + # ================================================================ # Session 管理 # ================================================================ @@ -439,12 +448,14 @@ class ChatService(MessagePersistenceMixin): session_id=req.session_id, user_id=user_id, request_type="v2_chat", - model=req.model or "", + model=self._metrics_model_name(req.model), ) req_metrics = RequestMetrics( session_id=req.session_id or "", user_id=user_id, request_type="v2_chat", + project_id=str(req.context.get("project_id") or ""), + model=self._metrics_model_name(req.model), ) main_collector = MessagePartsCollector() @@ -659,12 +670,14 @@ class ChatService(MessagePersistenceMixin): session_id=req.session_id, user_id=user_id, request_type="v2_chat", - model=req.model or "", + model=self._metrics_model_name(req.model), ) req_metrics = RequestMetrics( session_id=req.session_id, user_id=user_id, request_type="v2_chat", + project_id=str(req.context.get("project_id") or ""), + model=self._metrics_model_name(req.model), ) try: diff --git a/ocai-service/ocai/services/usage_metrics_service.py b/ocai-service/ocai/services/usage_metrics_service.py new file mode 100644 index 0000000..acf4511 --- /dev/null +++ b/ocai-service/ocai/services/usage_metrics_service.py @@ -0,0 +1,39 @@ +from __future__ import annotations + +from datetime import datetime +from typing import Any + +from ocai.repo.mysql.usage_repo import MysqlUsageRepository + + +class UsageMetricsService: + def __init__(self, repository: MysqlUsageRepository): + self._repository = repository + + def query( + self, + *, + user_id: str, + project_id: str | None = None, + model: str | None = None, + start: datetime | None = None, + end: datetime | None = None, + group_by: str = "model", + ) -> dict[str, Any]: + groups = self._repository.aggregate( + user_id=user_id, + project_id=project_id, + model=model, + start=start, + end=end, + group_by=group_by, + ) + return { + "groups": groups, + "totals": { + "requests": sum(item["requests"] for item in groups), + "input_tokens": sum(item["input_tokens"] for item in groups), + "output_tokens": sum(item["output_tokens"] for item in groups), + "estimated_cost": round(sum(item["estimated_cost"] for item in groups), 6), + }, + } diff --git a/ocai-service/sql/0002_usage_records.down.sql b/ocai-service/sql/0002_usage_records.down.sql new file mode 100644 index 0000000..c2817ba --- /dev/null +++ b/ocai-service/sql/0002_usage_records.down.sql @@ -0,0 +1 @@ +DROP TABLE IF EXISTS usage_records; diff --git a/ocai-service/sql/0002_usage_records.up.sql b/ocai-service/sql/0002_usage_records.up.sql new file mode 100644 index 0000000..2c06f4e --- /dev/null +++ b/ocai-service/sql/0002_usage_records.up.sql @@ -0,0 +1,16 @@ +CREATE TABLE IF NOT EXISTS usage_records ( + id VARCHAR(64) NOT NULL, + user_id VARCHAR(128) NOT NULL, + session_id VARCHAR(64) NOT NULL, + project_id VARCHAR(128) NULL, + model VARCHAR(128) NOT NULL, + input_tokens INTEGER NOT NULL, + output_tokens INTEGER NOT NULL, + duration_ms FLOAT NOT NULL, + estimated_cost FLOAT NOT NULL, + created_at DATETIME NOT NULL, + PRIMARY KEY (id), + INDEX idx_usage_user_created (user_id, created_at), + INDEX idx_usage_project_created (project_id, created_at), + INDEX idx_usage_model_created (model, created_at) +) ENGINE=InnoDB DEFAULT CHARSET=utf8mb4 COLLATE=utf8mb4_unicode_ci; diff --git a/ocai-service/tests/ai/test_metrics_persistence.py b/ocai-service/tests/ai/test_metrics_persistence.py new file mode 100644 index 0000000..8e95a0b --- /dev/null +++ b/ocai-service/tests/ai/test_metrics_persistence.py @@ -0,0 +1,36 @@ +import threading + +from ocai.ai.metrics import MetricsCollector, RequestMetrics + + +def test_metrics_are_persisted_asynchronously_once(): + saved = [] + persisted = threading.Event() + + def save(snapshot): + saved.append(snapshot) + persisted.set() + + collector = MetricsCollector( + persist_fn=save, + model_pricing={"qwen": (2.0, 8.0)}, + ) + metrics = RequestMetrics( + session_id="s1", + user_id="alice", + project_id="42", + model="qwen", + ) + metrics.on_token_usage(input_tokens=10, output_tokens=4) + + collector.record(metrics) + collector.record(metrics) + + assert persisted.wait(1) + collector.shutdown() + assert len(saved) == 1 + assert saved[0]["project_id"] == "42" + assert saved[0]["model"] == "qwen" + assert saved[0]["input_tokens"] == 10 + assert saved[0]["output_tokens"] == 4 + assert saved[0]["estimated_cost"] == 0.000052 diff --git a/ocai-service/tests/api/test_usage_metrics_endpoint.py b/ocai-service/tests/api/test_usage_metrics_endpoint.py new file mode 100644 index 0000000..7612e53 --- /dev/null +++ b/ocai-service/tests/api/test_usage_metrics_endpoint.py @@ -0,0 +1,34 @@ +from fastapi import FastAPI +from fastapi.testclient import TestClient + +from ocai.api.v2_route import router +from ocai.permission import UserInfo, get_current_user + + +class _UsageService: + def __init__(self): + self.kwargs = None + + def query(self, **kwargs): + self.kwargs = kwargs + return {"groups": [], "totals": {"requests": 0}} + + +def test_usage_query_is_scoped_to_current_user(): + app = FastAPI() + service = _UsageService() + app.state.usage_metrics_service = service + app.include_router(router) + app.dependency_overrides[get_current_user] = lambda: UserInfo(login_name="alice") + client = TestClient(app) + + response = client.get( + "/api/v2/ocai/metrics/usage", + params={"project_id": "42", "model": "qwen", "group_by": "month"}, + ) + + assert response.status_code == 200 + assert service.kwargs["user_id"] == "alice" + assert service.kwargs["project_id"] == "42" + assert service.kwargs["model"] == "qwen" + assert service.kwargs["group_by"] == "month" diff --git a/ocai-service/tests/repo/test_mysql_repos.py b/ocai-service/tests/repo/test_mysql_repos.py index 2f9b6f7..e84722e 100644 --- a/ocai-service/tests/repo/test_mysql_repos.py +++ b/ocai-service/tests/repo/test_mysql_repos.py @@ -20,6 +20,7 @@ from ocai.repo.mysql import ( MysqlAgentRepository, MysqlChatSessionRepository, MySQLEngine, + MysqlUsageRepository, MysqlUserPreferencesRepo, MysqlV2SessionRepository, ) @@ -129,6 +130,52 @@ def test_user_preferences_persist(engine): assert got["enabled_mcp_tools"] == ["x"] +def test_usage_records_are_persisted_and_aggregated(engine): + repo = MysqlUsageRepository(engine) + repo.save_request_metrics( + { + "user_id": "alice", + "session_id": "s1", + "project_id": "42", + "model": "qwen", + "input_tokens": 100, + "output_tokens": 40, + "total_duration_ms": 1200, + "estimated_cost": 0.012, + "timestamp": 1_787_500_000, + } + ) + repo.save_request_metrics( + { + "user_id": "alice", + "session_id": "s2", + "project_id": "42", + "model": "qwen", + "input_tokens": 50, + "output_tokens": 10, + "total_duration_ms": 800, + "estimated_cost": 0.004, + "timestamp": 1_787_500_100, + } + ) + + groups = repo.aggregate(user_id="alice", project_id="42", group_by="model") + monthly = repo.aggregate(user_id="alice", project_id="42", group_by="month") + + assert groups == [ + { + "group": "qwen", + "requests": 2, + "input_tokens": 150, + "output_tokens": 50, + "estimated_cost": 0.016, + "avg_duration_ms": 1000.0, + } + ] + assert monthly[0]["group"] == "2026-08" + assert monthly[0]["requests"] == 2 + + # --------------------------------------------------------------------------- # Agent logs + memory # --------------------------------------------------------------------------- -- Gitee