mirror of
https://github.com/xinnan-tech/xiaozhi-esp32-server.git
synced 2026-07-22 23:23:55 +08:00
feat: add FastAPI manager API compatibility baseline
This commit is contained in:
@@ -0,0 +1,20 @@
|
||||
APP_ENVIRONMENT=development
|
||||
APP_HOST=0.0.0.0
|
||||
APP_PORT=8002
|
||||
APP_CONTEXT_PATH=/xiaozhi
|
||||
APP_TIMEZONE=Asia/Shanghai
|
||||
APP_DATABASE_URL=mysql+asyncmy://xiaozhi:replace-me@mysql:3306/xiaozhi_esp32_server?charset=utf8mb4
|
||||
APP_REDIS_URL=redis://redis:6379/0
|
||||
# Local default; the container image overrides this with /data/uploads.
|
||||
APP_UPLOAD_DIR=./uploadfile
|
||||
# Docker Compose source: use a named volume by default, or set an existing
|
||||
# Java uploadfile host path while the implementations coexist.
|
||||
MANAGER_API_UPLOAD_SOURCE=manager-api-uploads
|
||||
APP_JAVA_RESOURCES_DIR=/opt/xiaozhi/java-resources
|
||||
APP_EXTERNAL_REQUEST_TIMEOUT_SECONDS=10
|
||||
APP_TRUSTED_PROXY_COUNT=1
|
||||
APP_LOG_LEVEL=INFO
|
||||
APP_GRACEFUL_SHUTDOWN_SECONDS=30
|
||||
# Test-only escape hatches. Leave both unset in deployments.
|
||||
# APP_SERVER_SECRET_OVERRIDE=
|
||||
# APP_ALLOW_START_WITHOUT_DEPENDENCIES=false
|
||||
@@ -0,0 +1,16 @@
|
||||
.env
|
||||
.mypy_cache/
|
||||
.pytest_cache/
|
||||
.ruff_cache/
|
||||
.test-runtime/
|
||||
.venv/
|
||||
.coverage
|
||||
coverage.xml
|
||||
htmlcov/
|
||||
__pycache__/
|
||||
*.py[cod]
|
||||
data/uploads/
|
||||
target/
|
||||
dist/
|
||||
!compatibility/*.json
|
||||
!tests/fixtures/*.json
|
||||
@@ -0,0 +1 @@
|
||||
3.10
|
||||
@@ -0,0 +1,54 @@
|
||||
FROM python:3.10.20-bookworm AS build
|
||||
|
||||
ENV UV_COMPILE_BYTECODE=1 \
|
||||
UV_LINK_MODE=copy \
|
||||
UV_HTTP_TIMEOUT=120 \
|
||||
UV_HTTP_RETRIES=10 \
|
||||
PATH=/app/.venv/bin:$PATH
|
||||
|
||||
WORKDIR /app
|
||||
|
||||
COPY --from=ghcr.io/astral-sh/uv:0.11.28 /uv /uvx /bin/
|
||||
|
||||
COPY main/manager-api-fastapi/pyproject.toml main/manager-api-fastapi/uv.lock main/manager-api-fastapi/README.md ./
|
||||
COPY main/manager-api-fastapi/app ./app
|
||||
RUN --mount=type=cache,target=/root/.cache/uv \
|
||||
uv sync --frozen --no-dev --no-editable
|
||||
|
||||
FROM python:3.10.20-slim-bookworm
|
||||
|
||||
ENV PYTHONDONTWRITEBYTECODE=1 \
|
||||
PYTHONUNBUFFERED=1 \
|
||||
PATH=/app/.venv/bin:$PATH \
|
||||
APP_JAVA_RESOURCES_DIR=/opt/xiaozhi/java-resources \
|
||||
APP_UPLOAD_DIR=/data/uploads \
|
||||
APP_HOST=0.0.0.0 \
|
||||
APP_PORT=8002 \
|
||||
APP_TIMEZONE=Asia/Shanghai
|
||||
|
||||
WORKDIR /app
|
||||
|
||||
RUN groupadd --gid 10001 xiaozhi \
|
||||
&& useradd --uid 10001 --gid xiaozhi --create-home --shell /usr/sbin/nologin xiaozhi
|
||||
|
||||
COPY --from=build --chown=10001:10001 /app/.venv ./.venv
|
||||
COPY --from=build --chown=10001:10001 /app/app ./app
|
||||
|
||||
COPY main/manager-api-fastapi/scripts/container-entrypoint.sh /usr/local/bin/manager-api-entrypoint
|
||||
COPY main/manager-api/src/main/resources/i18n /opt/xiaozhi/java-resources/i18n
|
||||
COPY main/manager-api/src/main/resources/db /opt/xiaozhi/java-resources/db
|
||||
|
||||
RUN mkdir -p /data/uploads \
|
||||
&& ln -s /data/uploads /app/uploadfile \
|
||||
&& chown -R xiaozhi:xiaozhi /data/uploads /opt/xiaozhi \
|
||||
&& chmod 0555 /usr/local/bin/manager-api-entrypoint
|
||||
|
||||
USER 10001:10001
|
||||
EXPOSE 8002
|
||||
VOLUME ["/data/uploads"]
|
||||
STOPSIGNAL SIGTERM
|
||||
|
||||
HEALTHCHECK --interval=15s --timeout=3s --start-period=20s --retries=4 \
|
||||
CMD ["python", "-c", "import urllib.request; urllib.request.urlopen('http://127.0.0.1:8002/xiaozhi/health/live', timeout=2).read()"]
|
||||
|
||||
ENTRYPOINT ["/usr/local/bin/manager-api-entrypoint"]
|
||||
@@ -0,0 +1,35 @@
|
||||
FROM maven:3.9.9-eclipse-temurin-21 AS build
|
||||
|
||||
WORKDIR /migration
|
||||
COPY main/manager-api-fastapi/migration-pom.xml ./pom.xml
|
||||
COPY main/manager-api-fastapi/migration-src ./migration-src
|
||||
COPY main/manager-api/src/main/resources ./java-resources
|
||||
# Keep the Maven repository outside the committed layer so an interrupted
|
||||
# registry transfer can resume on the next build. Resolver downloads are
|
||||
# deliberately serial: Apple Container's BuildKit NAT has proved unreliable
|
||||
# when several Maven Central responses are multiplexed over one connection.
|
||||
RUN --mount=type=cache,target=/root/.m2/repository \
|
||||
mvn -B \
|
||||
-Dmaven.repo.local=/root/.m2/repository \
|
||||
-Djava.resources.dir=/migration/java-resources \
|
||||
-Daether.connector.basic.threads=1 \
|
||||
-Daether.connector.connectTimeout=15000 \
|
||||
-Daether.connector.requestTimeout=60000 \
|
||||
-Daether.connector.http.retryHandler.count=5 \
|
||||
-Daether.connector.http.retryHandler.interval=1000 \
|
||||
-Daether.connector.http.retryHandler.intervalMax=5000 \
|
||||
package
|
||||
|
||||
FROM eclipse-temurin:21-jre
|
||||
WORKDIR /migration
|
||||
COPY --from=build /migration/target/manager-api-liquibase-runner-1.0.0-all.jar ./runner.jar
|
||||
COPY main/manager-api-fastapi/scripts/run-migrations.sh /usr/local/bin/run-manager-api-migrations
|
||||
RUN groupadd --gid 10001 xiaozhi \
|
||||
&& useradd --uid 10001 --gid xiaozhi --create-home --shell /usr/sbin/nologin xiaozhi \
|
||||
&& chown -R xiaozhi:xiaozhi /migration \
|
||||
&& chmod 0555 /usr/local/bin/run-manager-api-migrations
|
||||
ENV MIGRATION_RUNNER_JAR=/migration/runner.jar \
|
||||
TZ=Asia/Shanghai
|
||||
USER 10001:10001
|
||||
STOPSIGNAL SIGTERM
|
||||
ENTRYPOINT ["/usr/local/bin/run-manager-api-migrations"]
|
||||
@@ -0,0 +1,14 @@
|
||||
FROM nginx:1.28.0-alpine
|
||||
|
||||
COPY main/manager-api-fastapi/deploy/nginx.conf /etc/nginx/nginx.conf.template
|
||||
COPY main/manager-api-fastapi/deploy/nginx-entrypoint.sh /usr/local/bin/manager-api-nginx-entrypoint
|
||||
RUN MANAGER_API_UPSTREAM=127.0.0.1:8002 \
|
||||
envsubst '${MANAGER_API_UPSTREAM}' \
|
||||
< /etc/nginx/nginx.conf.template \
|
||||
> /tmp/nginx-build-check.conf \
|
||||
&& nginx -t -c /tmp/nginx-build-check.conf \
|
||||
&& rm /tmp/nginx-build-check.conf \
|
||||
&& chmod 0555 /usr/local/bin/manager-api-nginx-entrypoint
|
||||
|
||||
ENV MANAGER_API_UPSTREAM=manager-api-fastapi:8002
|
||||
ENTRYPOINT ["/usr/local/bin/manager-api-nginx-entrypoint"]
|
||||
@@ -0,0 +1,38 @@
|
||||
# manager-api-fastapi
|
||||
|
||||
`manager-api-fastapi` is the Python/FastAPI implementation of the existing Spring Boot
|
||||
`main/manager-api`. The Java service remains in the repository as the contract baseline,
|
||||
Liquibase migration owner, and rollback implementation.
|
||||
|
||||
## Local development
|
||||
|
||||
The service requires Python 3.10, MySQL 8, and Redis 5 or newer. Never point tests at a
|
||||
development database: the integration harness creates a dedicated database and Redis
|
||||
namespace/instance.
|
||||
|
||||
```bash
|
||||
cd main/manager-api-fastapi
|
||||
cp .env.example .env
|
||||
uv sync --locked
|
||||
uv run python -m app
|
||||
```
|
||||
|
||||
The compatible base URL is `http://127.0.0.1:8002/xiaozhi`. OpenAPI is exposed at
|
||||
`/xiaozhi/v3/api-docs` and the Swagger UI at `/xiaozhi/doc.html`.
|
||||
|
||||
Production must set `APP_DATABASE_URL`, `APP_REDIS_URL`, `APP_UPLOAD_DIR`, and
|
||||
`APP_JAVA_RESOURCES_DIR`. The last path must contain the original Java i18n resources and
|
||||
Liquibase changelog. `APP_SERVER_SECRET_OVERRIDE` is reserved for isolated tests; leaving it
|
||||
set in a deployment bypasses the database-backed `server.secret` lookup and is unsupported.
|
||||
|
||||
## Commands
|
||||
|
||||
```bash
|
||||
uv run pytest
|
||||
uv run ruff check app tests scripts
|
||||
uv run mypy app
|
||||
uv run python scripts/extract_java_routes.py --output compatibility/java-routes.json
|
||||
```
|
||||
|
||||
Migration, container, differential-contract, and cutover instructions are maintained in the
|
||||
repository-level migration documents under `docs/manager-api-fastapi-*.md`.
|
||||
@@ -0,0 +1 @@
|
||||
"""Xiaozhi manager API FastAPI implementation."""
|
||||
@@ -0,0 +1,7 @@
|
||||
import uvicorn
|
||||
|
||||
from app.core.config import get_settings
|
||||
|
||||
if __name__ == "__main__":
|
||||
settings = get_settings()
|
||||
uvicorn.run("app.main:app", host=settings.host, port=settings.port, log_level=settings.log_level.lower())
|
||||
@@ -0,0 +1 @@
|
||||
"""Shared compatibility infrastructure."""
|
||||
@@ -0,0 +1,72 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from functools import lru_cache
|
||||
from pathlib import Path
|
||||
from typing import Literal
|
||||
|
||||
from pydantic import Field, field_validator
|
||||
from pydantic_settings import BaseSettings, SettingsConfigDict
|
||||
|
||||
|
||||
def _default_java_resources() -> Path:
|
||||
return Path(__file__).resolve().parents[3] / "manager-api" / "src" / "main" / "resources"
|
||||
|
||||
|
||||
class Settings(BaseSettings):
|
||||
model_config = SettingsConfigDict(
|
||||
env_file=".env",
|
||||
env_prefix="APP_",
|
||||
case_sensitive=False,
|
||||
extra="ignore",
|
||||
)
|
||||
|
||||
environment: Literal["development", "test", "production"] = "development"
|
||||
host: str = "0.0.0.0" # noqa: S104 - container bind is intentional
|
||||
port: int = 8002
|
||||
context_path: str = "/xiaozhi"
|
||||
timezone: str = "Asia/Shanghai"
|
||||
database_url: str = "mysql+asyncmy://root:change-me@127.0.0.1:3306/xiaozhi_esp32_server?charset=utf8mb4"
|
||||
redis_url: str = "redis://127.0.0.1:6379/0"
|
||||
upload_dir: Path = Path("uploadfile")
|
||||
java_resources_dir: Path = Field(default_factory=_default_java_resources)
|
||||
external_request_timeout_seconds: float = 10.0
|
||||
database_pool_size: int = 20
|
||||
database_max_overflow: int = 20
|
||||
trusted_proxy_count: int = 1
|
||||
log_level: str = "INFO"
|
||||
server_secret_override: str | None = None
|
||||
allow_start_without_dependencies: bool = False
|
||||
job_lock_ttl_seconds: int = 120
|
||||
graceful_shutdown_seconds: float = 30.0
|
||||
|
||||
@field_validator("context_path")
|
||||
@classmethod
|
||||
def normalize_context_path(cls, value: str) -> str:
|
||||
normalized = "/" + value.strip("/")
|
||||
return "" if normalized == "/" else normalized
|
||||
|
||||
@field_validator("database_url")
|
||||
@classmethod
|
||||
def require_async_driver(cls, value: str) -> str:
|
||||
if value.startswith("mysql://"):
|
||||
return value.replace("mysql://", "mysql+asyncmy://", 1)
|
||||
if value.startswith("sqlite:///"):
|
||||
return value.replace("sqlite:///", "sqlite+aiosqlite:///", 1)
|
||||
return value
|
||||
|
||||
@property
|
||||
def i18n_dir(self) -> Path:
|
||||
return self.java_resources_dir / "i18n"
|
||||
|
||||
@property
|
||||
def changelog_path(self) -> Path:
|
||||
return self.java_resources_dir / "db" / "changelog" / "db.changelog-master.yaml"
|
||||
|
||||
|
||||
@lru_cache(maxsize=1)
|
||||
def get_settings() -> Settings:
|
||||
return Settings()
|
||||
|
||||
|
||||
def clear_settings_cache() -> None:
|
||||
get_settings.cache_clear()
|
||||
@@ -0,0 +1,63 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import hashlib
|
||||
import secrets
|
||||
import uuid
|
||||
|
||||
import bcrypt
|
||||
from gmssl import func, sm2 # type: ignore[import-untyped]
|
||||
|
||||
|
||||
def generate_database_token(value: str | None = None) -> str:
|
||||
source = value if value is not None else str(uuid.uuid4())
|
||||
return hashlib.md5(source.encode("utf-8"), usedforsecurity=False).hexdigest()
|
||||
|
||||
|
||||
def bcrypt_hash(password: str, rounds: int = 10) -> str:
|
||||
encoded = bcrypt.hashpw(password.encode("utf-8"), bcrypt.gensalt(rounds=rounds))
|
||||
# The bundled Java BCryptPasswordEncoder only accepts $2a$ hashes.
|
||||
return encoded.decode("ascii").replace("$2b$", "$2a$", 1)
|
||||
|
||||
|
||||
def bcrypt_matches(password: str, encoded: str | None) -> bool:
|
||||
if not encoded or not encoded.startswith(("$2a$", "$2$")):
|
||||
return False
|
||||
normalized = encoded.replace("$2$", "$2a$", 1)
|
||||
try:
|
||||
return bcrypt.checkpw(password.encode("utf-8"), normalized.encode("ascii"))
|
||||
except (ValueError, UnicodeEncodeError):
|
||||
return False
|
||||
|
||||
|
||||
def sm2_generate_keypair() -> tuple[str, str]:
|
||||
private_key = func.random_hex(64)
|
||||
helper = sm2.CryptSM2(private_key=private_key, public_key="", mode=1)
|
||||
public_point = str(helper._kg(int(private_key, 16), sm2.default_ecc_table["g"])) # noqa: SLF001
|
||||
public_key = "04" + public_point
|
||||
return public_key, private_key
|
||||
|
||||
|
||||
def sm2_encrypt_c1c3c2(public_key: str, plaintext: str) -> str:
|
||||
helper = sm2.CryptSM2(private_key="", public_key=public_key, mode=1)
|
||||
encrypted = helper.encrypt(plaintext.encode("utf-8"))
|
||||
if encrypted is None:
|
||||
raise ValueError("SM2 KDF returned an all-zero key")
|
||||
# BouncyCastle's SM2Engine emits the uncompressed-point marker.
|
||||
return "04" + bytes(encrypted).hex()
|
||||
|
||||
|
||||
def sm2_decrypt_c1c3c2(private_key: str, ciphertext: str) -> str:
|
||||
normalized = ciphertext.strip().lower()
|
||||
if normalized.startswith("04"):
|
||||
normalized = normalized[2:]
|
||||
if len(normalized) < 128 + 64 or len(normalized) % 2:
|
||||
raise ValueError("invalid SM2 C1C3C2 ciphertext")
|
||||
helper = sm2.CryptSM2(private_key=private_key, public_key="", mode=1)
|
||||
decrypted = helper.decrypt(bytes.fromhex(normalized))
|
||||
if decrypted is None:
|
||||
raise ValueError("SM2 decryption failed")
|
||||
return bytes(decrypted).decode("utf-8")
|
||||
|
||||
|
||||
def random_hex(length: int) -> str:
|
||||
return secrets.token_hex((length + 1) // 2)[:length]
|
||||
@@ -0,0 +1,96 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from collections.abc import AsyncIterator, Mapping, Sequence
|
||||
from contextlib import asynccontextmanager
|
||||
from typing import Any
|
||||
|
||||
from sqlalchemy import Result, TextClause, text
|
||||
from sqlalchemy.ext.asyncio import AsyncEngine, AsyncSession, async_sessionmaker, create_async_engine
|
||||
|
||||
from app.core.config import Settings, get_settings
|
||||
|
||||
_engine: AsyncEngine | None = None
|
||||
_session_factory: async_sessionmaker[AsyncSession] | None = None
|
||||
|
||||
|
||||
def configure_database(settings: Settings | None = None) -> tuple[AsyncEngine, async_sessionmaker[AsyncSession]]:
|
||||
global _engine, _session_factory
|
||||
selected = settings or get_settings()
|
||||
engine_options: dict[str, Any] = {"pool_pre_ping": True}
|
||||
if not selected.database_url.startswith("sqlite"):
|
||||
engine_options.update(pool_size=selected.database_pool_size, max_overflow=selected.database_max_overflow)
|
||||
_engine = create_async_engine(selected.database_url, **engine_options)
|
||||
_session_factory = async_sessionmaker(_engine, expire_on_commit=False, autoflush=False)
|
||||
return _engine, _session_factory
|
||||
|
||||
|
||||
def get_engine() -> AsyncEngine:
|
||||
if _engine is None:
|
||||
return configure_database()[0]
|
||||
return _engine
|
||||
|
||||
|
||||
def get_session_factory() -> async_sessionmaker[AsyncSession]:
|
||||
if _session_factory is None:
|
||||
return configure_database()[1]
|
||||
return _session_factory
|
||||
|
||||
|
||||
async def get_db() -> AsyncIterator[AsyncSession]:
|
||||
async with get_session_factory()() as session:
|
||||
yield session
|
||||
|
||||
|
||||
@asynccontextmanager
|
||||
async def transaction() -> AsyncIterator[AsyncSession]:
|
||||
async with get_session_factory()() as session, session.begin():
|
||||
yield session
|
||||
|
||||
|
||||
async def dispose_database() -> None:
|
||||
global _engine, _session_factory
|
||||
if _engine is not None:
|
||||
await _engine.dispose()
|
||||
_engine = None
|
||||
_session_factory = None
|
||||
|
||||
|
||||
class Repository:
|
||||
def __init__(self, session: AsyncSession):
|
||||
self.session = session
|
||||
|
||||
@staticmethod
|
||||
def statement(sql: str | TextClause) -> TextClause:
|
||||
return text(sql) if isinstance(sql, str) else sql
|
||||
|
||||
async def fetch_one(self, sql: str | TextClause, params: Mapping[str, Any] | None = None) -> dict[str, Any] | None:
|
||||
result = await self.session.execute(self.statement(sql), dict(params or {}))
|
||||
row = result.mappings().first()
|
||||
return dict(row) if row is not None else None
|
||||
|
||||
async def fetch_all(self, sql: str | TextClause, params: Mapping[str, Any] | None = None) -> list[dict[str, Any]]:
|
||||
result = await self.session.execute(self.statement(sql), dict(params or {}))
|
||||
return [dict(row) for row in result.mappings().all()]
|
||||
|
||||
async def scalar(self, sql: str | TextClause, params: Mapping[str, Any] | None = None) -> Any:
|
||||
result = await self.session.execute(self.statement(sql), dict(params or {}))
|
||||
return result.scalar_one_or_none()
|
||||
|
||||
async def execute(self, sql: str | TextClause, params: Mapping[str, Any] | None = None) -> int:
|
||||
result: Result[Any] = await self.session.execute(self.statement(sql), dict(params or {}))
|
||||
return int(getattr(result, "rowcount", 0) or 0)
|
||||
|
||||
async def execute_many(self, sql: str | TextClause, params: Sequence[Mapping[str, Any]]) -> int:
|
||||
if not params:
|
||||
return 0
|
||||
result: Result[Any] = await self.session.execute(self.statement(sql), [dict(item) for item in params])
|
||||
return int(getattr(result, "rowcount", 0) or 0)
|
||||
|
||||
|
||||
async def database_ping() -> bool:
|
||||
try:
|
||||
async with get_session_factory()() as session:
|
||||
await session.execute(text("SELECT 1"))
|
||||
return True
|
||||
except Exception:
|
||||
return False
|
||||
@@ -0,0 +1,44 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from dataclasses import dataclass
|
||||
|
||||
|
||||
class ErrorCode:
|
||||
INTERNAL_SERVER_ERROR = 500
|
||||
UNAUTHORIZED = 401
|
||||
FORBIDDEN = 403
|
||||
DB_RECORD_EXISTS = 10002
|
||||
PARAMS_GET_ERROR = 10003
|
||||
ACCOUNT_PASSWORD_ERROR = 10004
|
||||
ACCOUNT_DISABLE = 10005
|
||||
CAPTCHA_ERROR = 10007
|
||||
PASSWORD_ERROR = 10009
|
||||
UPLOAD_FILE_EMPTY = 10019
|
||||
TOKEN_INVALID = 10021
|
||||
ACCOUNT_LOCK = 10022
|
||||
INVALID_SYMBOL = 10029
|
||||
PASSWORD_LENGTH_ERROR = 10030
|
||||
PASSWORD_WEAK_ERROR = 10031
|
||||
DEL_MYSELF_ERROR = 10032
|
||||
DEVICE_CAPTCHA_ERROR = 10033
|
||||
PARAM_VALUE_NULL = 10034
|
||||
PARAM_TYPE_NULL = 10035
|
||||
PARAM_TYPE_INVALID = 10036
|
||||
PARAM_NUMBER_INVALID = 10037
|
||||
PARAM_BOOLEAN_INVALID = 10038
|
||||
PARAM_ARRAY_INVALID = 10039
|
||||
PARAM_JSON_INVALID = 10040
|
||||
RESOURCE_NOT_FOUND = 10051
|
||||
ADD_DATA_FAILED = 10065
|
||||
UPDATE_DATA_FAILED = 10066
|
||||
MODEL_TYPE_PROVIDE_CODE_NOT_NULL = 10131
|
||||
|
||||
|
||||
@dataclass(slots=True)
|
||||
class AppError(Exception):
|
||||
code: int
|
||||
message: str | None = None
|
||||
params: tuple[object, ...] = ()
|
||||
|
||||
def __str__(self) -> str:
|
||||
return self.message or str(self.code)
|
||||
@@ -0,0 +1,90 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import re
|
||||
from functools import lru_cache
|
||||
from pathlib import Path
|
||||
|
||||
from app.core.config import get_settings
|
||||
|
||||
LANGUAGE_FILES: dict[str, str] = {
|
||||
"zh-CN": "messages_zh_CN.properties",
|
||||
"zh-TW": "messages_zh_TW.properties",
|
||||
"en-US": "messages_en_US.properties",
|
||||
"de-DE": "messages_de_DE.properties",
|
||||
"vi-VN": "messages_vi_VN.properties",
|
||||
"pt-BR": "messages_pt_BR.properties",
|
||||
}
|
||||
_UNICODE_ESCAPE = re.compile(r"\\u([0-9a-fA-F]{4})")
|
||||
|
||||
|
||||
def resolve_language(accept_language: str | None) -> str:
|
||||
if not accept_language:
|
||||
return "zh-CN"
|
||||
primary = accept_language.split(",", 1)[0].split(";", 1)[0].strip().replace("_", "-")
|
||||
exact = {key.lower(): key for key in LANGUAGE_FILES}
|
||||
if primary.lower() in exact:
|
||||
return exact[primary.lower()]
|
||||
prefix = primary.lower().split("-", 1)[0]
|
||||
return {
|
||||
"zh": "zh-CN",
|
||||
"en": "en-US",
|
||||
"de": "de-DE",
|
||||
"vi": "vi-VN",
|
||||
"pt": "pt-BR",
|
||||
}.get(prefix, "zh-CN")
|
||||
|
||||
|
||||
def _unescape(value: str) -> str:
|
||||
decoded = _UNICODE_ESCAPE.sub(lambda match: chr(int(match.group(1), 16)), value)
|
||||
return (
|
||||
decoded.replace("\\t", "\t")
|
||||
.replace("\\n", "\n")
|
||||
.replace("\\r", "\r")
|
||||
.replace("\\f", "\f")
|
||||
.replace("\\=", "=")
|
||||
.replace("\\:", ":")
|
||||
.replace("\\ ", " ")
|
||||
.replace("\\\\", "\\")
|
||||
)
|
||||
|
||||
|
||||
def _load_properties(path: Path) -> dict[str, str]:
|
||||
messages: dict[str, str] = {}
|
||||
if not path.exists():
|
||||
return messages
|
||||
continuation = ""
|
||||
for raw_line in path.read_text(encoding="utf-8").splitlines():
|
||||
line = continuation + raw_line
|
||||
if line.endswith("\\") and not line.endswith("\\\\"):
|
||||
continuation = line[:-1]
|
||||
continue
|
||||
continuation = ""
|
||||
stripped = line.strip()
|
||||
if not stripped or stripped.startswith(("#", "!")):
|
||||
continue
|
||||
delimiter = "=" if "=" in line else ":"
|
||||
if delimiter not in line:
|
||||
continue
|
||||
key, value = line.split(delimiter, 1)
|
||||
messages[key.strip()] = _unescape(value.strip())
|
||||
return messages
|
||||
|
||||
|
||||
@lru_cache(maxsize=16)
|
||||
def messages_for(language: str, i18n_dir: str | None = None) -> dict[str, str]:
|
||||
directory = Path(i18n_dir) if i18n_dir else get_settings().i18n_dir
|
||||
default_messages = _load_properties(directory / "messages.properties")
|
||||
default_messages.update(_load_properties(directory / LANGUAGE_FILES.get(language, LANGUAGE_FILES["zh-CN"])))
|
||||
return default_messages
|
||||
|
||||
|
||||
def message_for(code: int, accept_language: str | None, *params: object) -> str:
|
||||
language = resolve_language(accept_language)
|
||||
template = messages_for(language).get(str(code), str(code))
|
||||
for index, param in enumerate(params):
|
||||
template = template.replace("{" + str(index) + "}", str(param))
|
||||
return template
|
||||
|
||||
|
||||
def clear_i18n_cache() -> None:
|
||||
messages_for.cache_clear()
|
||||
@@ -0,0 +1,56 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import os
|
||||
import socket
|
||||
import threading
|
||||
import time
|
||||
|
||||
|
||||
class SnowflakeIdGenerator:
|
||||
"""MyBatis-Plus compatible 41/5/5/12-bit Snowflake identifier generator."""
|
||||
|
||||
EPOCH = 1288834974657
|
||||
SEQUENCE_BITS = 12
|
||||
WORKER_BITS = 5
|
||||
DATACENTER_BITS = 5
|
||||
MAX_SEQUENCE = (1 << SEQUENCE_BITS) - 1
|
||||
WORKER_SHIFT = SEQUENCE_BITS
|
||||
DATACENTER_SHIFT = SEQUENCE_BITS + WORKER_BITS
|
||||
TIMESTAMP_SHIFT = SEQUENCE_BITS + WORKER_BITS + DATACENTER_BITS
|
||||
|
||||
def __init__(self, worker_id: int | None = None, datacenter_id: int | None = None):
|
||||
host_hash = sum(socket.gethostname().encode("utf-8"))
|
||||
self.worker_id = worker_id if worker_id is not None else (host_hash ^ os.getpid()) & 31
|
||||
self.datacenter_id = datacenter_id if datacenter_id is not None else host_hash & 31
|
||||
if not 0 <= self.worker_id <= 31 or not 0 <= self.datacenter_id <= 31:
|
||||
raise ValueError("worker_id and datacenter_id must be in [0, 31]")
|
||||
self._sequence = 0
|
||||
self._last_timestamp = -1
|
||||
self._lock = threading.Lock()
|
||||
|
||||
@staticmethod
|
||||
def _milliseconds() -> int:
|
||||
return time.time_ns() // 1_000_000
|
||||
|
||||
def next_id(self) -> int:
|
||||
with self._lock:
|
||||
timestamp = self._milliseconds()
|
||||
if timestamp < self._last_timestamp:
|
||||
raise RuntimeError("clock moved backwards; refusing to generate a duplicate Snowflake ID")
|
||||
if timestamp == self._last_timestamp:
|
||||
self._sequence = (self._sequence + 1) & self.MAX_SEQUENCE
|
||||
if self._sequence == 0:
|
||||
while timestamp <= self._last_timestamp:
|
||||
timestamp = self._milliseconds()
|
||||
else:
|
||||
self._sequence = 0
|
||||
self._last_timestamp = timestamp
|
||||
return (
|
||||
((timestamp - self.EPOCH) << self.TIMESTAMP_SHIFT)
|
||||
| (self.datacenter_id << self.DATACENTER_SHIFT)
|
||||
| (self.worker_id << self.WORKER_SHIFT)
|
||||
| self._sequence
|
||||
)
|
||||
|
||||
|
||||
snowflake = SnowflakeIdGenerator()
|
||||
@@ -0,0 +1,226 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import json
|
||||
import logging
|
||||
from collections.abc import AsyncIterator
|
||||
from contextlib import asynccontextmanager, suppress
|
||||
from datetime import datetime
|
||||
from typing import Any, cast
|
||||
from zoneinfo import ZoneInfo
|
||||
|
||||
from redis.asyncio import Redis
|
||||
|
||||
from app.core.config import get_settings
|
||||
|
||||
_client: Redis | None = None
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
class JavaRedisCodec:
|
||||
"""Wire-compatible subset of Spring Data's ``RedisSerializer.json()``.
|
||||
|
||||
Spring enables Jackson default typing for non-final values. Consequently a
|
||||
plain JSON map/list cannot be read by the retained Java rollback service. A
|
||||
map carries ``@class`` and a collection uses Jackson's wrapper-array form.
|
||||
``java_type`` and ``item_java_type`` cover the few caches whose Java readers
|
||||
cast values to concrete DTO/entity classes.
|
||||
"""
|
||||
|
||||
@staticmethod
|
||||
def encode(
|
||||
value: Any,
|
||||
*,
|
||||
java_type: str | None = None,
|
||||
item_java_type: str | None = None,
|
||||
field_java_types: dict[str, str] | None = None,
|
||||
) -> bytes:
|
||||
wire = JavaRedisCodec._encode_value(
|
||||
value,
|
||||
java_type=java_type,
|
||||
item_java_type=item_java_type,
|
||||
field_java_types=field_java_types,
|
||||
nested=False,
|
||||
)
|
||||
return json.dumps(wire, ensure_ascii=False, separators=(",", ":")).encode("utf-8")
|
||||
|
||||
@staticmethod
|
||||
def decode(value: bytes | str | None) -> Any:
|
||||
if value is None:
|
||||
return None
|
||||
raw = value.decode("utf-8") if isinstance(value, bytes) else value
|
||||
try:
|
||||
return JavaRedisCodec._decode_value(json.loads(raw))
|
||||
except json.JSONDecodeError:
|
||||
return raw
|
||||
|
||||
@staticmethod
|
||||
def _encode_value(
|
||||
value: Any,
|
||||
*,
|
||||
java_type: str | None = None,
|
||||
item_java_type: str | None = None,
|
||||
field_java_types: dict[str, str] | None = None,
|
||||
nested: bool = True,
|
||||
) -> Any:
|
||||
if value is None or isinstance(value, str | bool | float):
|
||||
return value
|
||||
if isinstance(value, int):
|
||||
# Jackson's default typing only adds the Long wrapper when the
|
||||
# runtime value sits behind an Object-typed container slot. A
|
||||
# top-level Long, or a field with a declared Long type, is emitted
|
||||
# as an ordinary JSON number.
|
||||
if not nested or java_type == "java.lang.Long" or -(2**31) <= value < 2**31:
|
||||
return value
|
||||
return ["java.lang.Long", value]
|
||||
if isinstance(value, datetime):
|
||||
timezone = ZoneInfo(get_settings().timezone)
|
||||
localized = value.replace(tzinfo=timezone) if value.tzinfo is None else value.astimezone(timezone)
|
||||
return ["java.util.Date", int(localized.timestamp() * 1000)]
|
||||
if isinstance(value, dict):
|
||||
selected_type = java_type or "java.util.HashMap"
|
||||
result: dict[str, Any] = {"@class": selected_type}
|
||||
pojo = selected_type not in {
|
||||
"java.util.HashMap",
|
||||
"java.util.LinkedHashMap",
|
||||
"java.util.TreeMap",
|
||||
"cn.hutool.json.JSONObject",
|
||||
}
|
||||
for raw_key, item in value.items():
|
||||
if raw_key == "@class":
|
||||
continue
|
||||
key = _snake_to_camel(str(raw_key)) if pojo else str(raw_key)
|
||||
child_type = (field_java_types or {}).get(key) or (field_java_types or {}).get(str(raw_key))
|
||||
result[key] = JavaRedisCodec._encode_value(item, java_type=child_type, nested=True)
|
||||
return result
|
||||
if isinstance(value, set | frozenset):
|
||||
return [
|
||||
"java.util.HashSet",
|
||||
[
|
||||
JavaRedisCodec._encode_value(item, java_type=item_java_type, nested=True)
|
||||
for item in value
|
||||
],
|
||||
]
|
||||
if isinstance(value, list | tuple):
|
||||
return [
|
||||
"java.util.ArrayList",
|
||||
[
|
||||
JavaRedisCodec._encode_value(item, java_type=item_java_type, nested=True)
|
||||
for item in value
|
||||
],
|
||||
]
|
||||
return value
|
||||
|
||||
@staticmethod
|
||||
def _decode_value(value: Any) -> Any:
|
||||
if isinstance(value, dict):
|
||||
return {
|
||||
str(key): JavaRedisCodec._decode_value(item)
|
||||
for key, item in value.items()
|
||||
if key != "@class"
|
||||
}
|
||||
if isinstance(value, list):
|
||||
if len(value) == 2 and isinstance(value[0], str) and value[0].startswith("java."):
|
||||
type_name, payload = value
|
||||
if type_name == "java.util.Date":
|
||||
timezone = ZoneInfo(get_settings().timezone)
|
||||
return datetime.fromtimestamp(float(payload) / 1000, timezone).replace(tzinfo=None)
|
||||
if type_name in {
|
||||
"java.util.ArrayList",
|
||||
"java.util.LinkedList",
|
||||
"java.util.HashSet",
|
||||
"java.util.LinkedHashSet",
|
||||
} and isinstance(payload, list):
|
||||
return [JavaRedisCodec._decode_value(item) for item in payload]
|
||||
return JavaRedisCodec._decode_value(payload)
|
||||
return [JavaRedisCodec._decode_value(item) for item in value]
|
||||
return value
|
||||
|
||||
|
||||
def _snake_to_camel(value: str) -> str:
|
||||
head, *tail = value.split("_")
|
||||
return head + "".join(part[:1].upper() + part[1:] for part in tail)
|
||||
|
||||
|
||||
def get_redis() -> Redis:
|
||||
global _client
|
||||
if _client is None:
|
||||
_client = Redis.from_url(get_settings().redis_url, decode_responses=False)
|
||||
return _client
|
||||
|
||||
|
||||
async def close_redis() -> None:
|
||||
global _client
|
||||
if _client is not None:
|
||||
await _client.aclose()
|
||||
_client = None
|
||||
|
||||
|
||||
async def redis_ping() -> bool:
|
||||
try:
|
||||
return bool(await get_redis().ping())
|
||||
except Exception:
|
||||
return False
|
||||
|
||||
|
||||
async def java_get(key: str) -> Any:
|
||||
return JavaRedisCodec.decode(await cast(Any, get_redis().get(key)))
|
||||
|
||||
|
||||
async def java_set(
|
||||
key: str,
|
||||
value: Any,
|
||||
ttl_seconds: int | None = None,
|
||||
*,
|
||||
java_type: str | None = None,
|
||||
item_java_type: str | None = None,
|
||||
) -> None:
|
||||
await cast(Any, get_redis().set)(
|
||||
key,
|
||||
JavaRedisCodec.encode(value, java_type=java_type, item_java_type=item_java_type),
|
||||
ex=ttl_seconds,
|
||||
)
|
||||
|
||||
|
||||
async def java_hget(key: str, field: str) -> Any:
|
||||
return JavaRedisCodec.decode(await cast(Any, get_redis().hget(key, field)))
|
||||
|
||||
|
||||
async def java_hset(key: str, field: str, value: Any, ttl_seconds: int = 86400) -> None:
|
||||
redis = get_redis()
|
||||
await cast(Any, redis.hset)(key, field, JavaRedisCodec.encode(value))
|
||||
await cast(Any, redis.expire)(key, ttl_seconds)
|
||||
|
||||
|
||||
@asynccontextmanager
|
||||
async def distributed_lock(name: str, ttl_seconds: int) -> AsyncIterator[bool]:
|
||||
lock = get_redis().lock(name, timeout=ttl_seconds, blocking_timeout=0)
|
||||
acquired = bool(await lock.acquire(blocking=False))
|
||||
renewal: asyncio.Task[None] | None = None
|
||||
if acquired:
|
||||
renewal = asyncio.create_task(_renew_lock(lock, ttl_seconds))
|
||||
try:
|
||||
yield acquired
|
||||
finally:
|
||||
if renewal is not None:
|
||||
renewal.cancel()
|
||||
with suppress(asyncio.CancelledError):
|
||||
await renewal
|
||||
if acquired:
|
||||
try:
|
||||
await lock.release()
|
||||
except Exception:
|
||||
logger.warning("Lost ownership of distributed lock %s before release", name, exc_info=True)
|
||||
|
||||
|
||||
async def _renew_lock(lock: Any, ttl_seconds: int) -> None:
|
||||
"""Keep a held job lock alive until its owner leaves the context."""
|
||||
|
||||
interval = max(float(ttl_seconds) / 3, 0.25)
|
||||
while True:
|
||||
await asyncio.sleep(interval)
|
||||
try:
|
||||
await lock.extend(ttl_seconds, replace_ttl=True)
|
||||
except Exception:
|
||||
logger.exception("Unable to renew distributed lock; duplicate execution protection is at risk")
|
||||
return
|
||||
@@ -0,0 +1,66 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
from typing import Any
|
||||
|
||||
from fastapi import Request
|
||||
from starlette.responses import JSONResponse, Response
|
||||
|
||||
from app.core.i18n import message_for
|
||||
from app.core.serialization import java_compatible
|
||||
|
||||
|
||||
class JavaJSONResponse(JSONResponse):
|
||||
def render(self, content: Any) -> bytes:
|
||||
return json.dumps(
|
||||
java_compatible(content),
|
||||
ensure_ascii=False,
|
||||
allow_nan=False,
|
||||
separators=(",", ":"),
|
||||
).encode("utf-8")
|
||||
|
||||
|
||||
def envelope(data: Any = None, *, code: int = 0, msg: str = "success") -> dict[str, Any]:
|
||||
return {"code": code, "msg": msg, "data": data}
|
||||
|
||||
|
||||
def ok(data: Any = None) -> JavaJSONResponse:
|
||||
return JavaJSONResponse(envelope(data))
|
||||
|
||||
|
||||
def error_response(
|
||||
request: Request,
|
||||
code: int,
|
||||
message: str | None = None,
|
||||
*,
|
||||
status_code: int = 200,
|
||||
params: tuple[object, ...] = (),
|
||||
media_type: str = "application/json",
|
||||
) -> JavaJSONResponse:
|
||||
translated = message or message_for(code, request.headers.get("Accept-Language"), *params)
|
||||
return JavaJSONResponse(
|
||||
envelope(None, code=code, msg=translated),
|
||||
status_code=status_code,
|
||||
media_type=media_type,
|
||||
)
|
||||
|
||||
|
||||
def raw_json(content: Any, *, exclude_none: bool = False, status_code: int = 200) -> Response:
|
||||
normalized = java_compatible(content)
|
||||
if exclude_none:
|
||||
normalized = _drop_none(normalized)
|
||||
body = json.dumps(normalized, ensure_ascii=False, allow_nan=False, separators=(",", ":")).encode("utf-8")
|
||||
return Response(
|
||||
body,
|
||||
status_code=status_code,
|
||||
media_type="application/json",
|
||||
headers={"Content-Length": str(len(body))},
|
||||
)
|
||||
|
||||
|
||||
def _drop_none(value: Any) -> Any:
|
||||
if isinstance(value, dict):
|
||||
return {key: _drop_none(item) for key, item in value.items() if item is not None}
|
||||
if isinstance(value, list):
|
||||
return [_drop_none(item) for item in value]
|
||||
return value
|
||||
@@ -0,0 +1,181 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import fnmatch
|
||||
import hmac
|
||||
from dataclasses import dataclass
|
||||
from datetime import datetime
|
||||
from typing import Any
|
||||
|
||||
from fastapi import Request
|
||||
from sqlalchemy import text
|
||||
from starlette.middleware.base import BaseHTTPMiddleware, RequestResponseEndpoint
|
||||
from starlette.responses import Response
|
||||
|
||||
from app.core.config import get_settings
|
||||
from app.core.database import get_session_factory
|
||||
from app.core.errors import AppError, ErrorCode
|
||||
from app.core.responses import error_response
|
||||
from app.services.system_params import SystemParamService
|
||||
|
||||
PUBLIC_PATTERNS = (
|
||||
"/ota/*",
|
||||
"/ota",
|
||||
"/otaMag/download/*",
|
||||
"/webjars/*",
|
||||
"/druid/*",
|
||||
"/v3/api-docs*",
|
||||
"/doc.html*",
|
||||
"/favicon.ico",
|
||||
"/user/captcha",
|
||||
"/user/smsVerification",
|
||||
"/user/login",
|
||||
"/user/pub-config",
|
||||
"/user/register",
|
||||
"/user/retrieve-password",
|
||||
"/api/ping",
|
||||
"/agent/chat-history/download/*",
|
||||
"/agent/play/*",
|
||||
"/voiceClone/play/*",
|
||||
"/health",
|
||||
"/health/live",
|
||||
"/health/ready",
|
||||
)
|
||||
SERVER_PATTERNS = (
|
||||
"/config/*",
|
||||
"/device/address-book/call",
|
||||
"/device/address-book/lookup",
|
||||
"/agent/chat-history/report",
|
||||
"/agent/chat-summary/*",
|
||||
"/agent/chat-title/*",
|
||||
)
|
||||
|
||||
|
||||
@dataclass(slots=True, frozen=True)
|
||||
class AuthUser:
|
||||
id: int
|
||||
username: str
|
||||
super_admin: int
|
||||
status: int
|
||||
token: str
|
||||
row: dict[str, Any]
|
||||
|
||||
@property
|
||||
def is_super_admin(self) -> bool:
|
||||
return self.super_admin == 1
|
||||
|
||||
|
||||
def _matches(path: str, patterns: tuple[str, ...]) -> bool:
|
||||
return any(fnmatch.fnmatchcase(path, pattern) for pattern in patterns)
|
||||
|
||||
|
||||
def _bearer_token(request: Request) -> str | None:
|
||||
authorization = request.headers.get("Authorization")
|
||||
if not authorization or not authorization.startswith("Bearer "):
|
||||
return None
|
||||
value = authorization[len("Bearer ") :]
|
||||
return value if value.strip() else None
|
||||
|
||||
|
||||
class AuthenticationMiddleware(BaseHTTPMiddleware):
|
||||
async def dispatch(self, request: Request, call_next: RequestResponseEndpoint) -> Response:
|
||||
if request.method == "OPTIONS":
|
||||
return await call_next(request)
|
||||
settings = get_settings()
|
||||
path = request.url.path
|
||||
if settings.context_path and path.startswith(settings.context_path):
|
||||
path = path[len(settings.context_path) :] or "/"
|
||||
if _matches(path, PUBLIC_PATTERNS):
|
||||
request.state.auth_mode = "anonymous"
|
||||
return await call_next(request)
|
||||
if _matches(path, SERVER_PATTERNS):
|
||||
return await self._server_auth(request, call_next)
|
||||
return await self._user_auth(request, call_next)
|
||||
|
||||
async def _server_auth(self, request: Request, call_next: RequestResponseEndpoint) -> Response:
|
||||
provided = _bearer_token(request)
|
||||
if provided is None:
|
||||
return error_response(
|
||||
request,
|
||||
ErrorCode.UNAUTHORIZED,
|
||||
"服务器密钥不能为空",
|
||||
media_type="application/json;charset=utf-8",
|
||||
)
|
||||
expected = get_settings().server_secret_override
|
||||
if expected is None:
|
||||
try:
|
||||
async with get_session_factory()() as session:
|
||||
expected = await SystemParamService(session).get_value("server.secret", from_cache=True)
|
||||
except Exception:
|
||||
expected = None
|
||||
if not expected or not hmac.compare_digest(provided, expected):
|
||||
return error_response(
|
||||
request,
|
||||
ErrorCode.UNAUTHORIZED,
|
||||
"无效的服务器密钥",
|
||||
media_type="application/json;charset=utf-8",
|
||||
)
|
||||
request.state.auth_mode = "server"
|
||||
return await call_next(request)
|
||||
|
||||
async def _user_auth(self, request: Request, call_next: RequestResponseEndpoint) -> Response:
|
||||
token = _bearer_token(request)
|
||||
if token is None:
|
||||
return error_response(
|
||||
request,
|
||||
ErrorCode.UNAUTHORIZED,
|
||||
media_type="application/json;charset=utf-8",
|
||||
)
|
||||
try:
|
||||
async with get_session_factory()() as session:
|
||||
result = await session.execute(
|
||||
text(
|
||||
"SELECT u.* FROM sys_user_token t "
|
||||
"JOIN sys_user u ON u.id = t.user_id "
|
||||
"WHERE t.token = :token AND t.expire_date >= CURRENT_TIMESTAMP LIMIT 1"
|
||||
),
|
||||
{"token": token},
|
||||
)
|
||||
mapping = result.mappings().first()
|
||||
except Exception:
|
||||
mapping = None
|
||||
if mapping is None or mapping.get("status") is None or int(mapping["status"]) != 1:
|
||||
return error_response(
|
||||
request,
|
||||
ErrorCode.UNAUTHORIZED,
|
||||
media_type="application/json;charset=utf-8",
|
||||
)
|
||||
row = dict(mapping)
|
||||
request.state.user = AuthUser(
|
||||
id=int(row["id"]),
|
||||
username=str(row.get("username") or ""),
|
||||
super_admin=int(row.get("super_admin") or 0),
|
||||
status=int(row["status"]),
|
||||
token=token,
|
||||
row=row,
|
||||
)
|
||||
request.state.auth_mode = "user"
|
||||
return await call_next(request)
|
||||
|
||||
|
||||
def current_user(request: Request) -> AuthUser:
|
||||
user = getattr(request.state, "user", None)
|
||||
if not isinstance(user, AuthUser):
|
||||
raise AppError(ErrorCode.UNAUTHORIZED)
|
||||
return user
|
||||
|
||||
|
||||
def require_normal(request: Request) -> AuthUser:
|
||||
return current_user(request)
|
||||
|
||||
|
||||
def require_super_admin(request: Request) -> AuthUser:
|
||||
user = current_user(request)
|
||||
if not user.is_super_admin:
|
||||
raise AppError(ErrorCode.FORBIDDEN)
|
||||
return user
|
||||
|
||||
|
||||
def shanghai_now_naive() -> datetime:
|
||||
from zoneinfo import ZoneInfo
|
||||
|
||||
return datetime.now(tz=ZoneInfo(get_settings().timezone)).replace(tzinfo=None)
|
||||
@@ -0,0 +1,105 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import dataclasses
|
||||
import re
|
||||
from collections.abc import Mapping, Sequence
|
||||
from datetime import date, datetime, time
|
||||
from decimal import Decimal
|
||||
from enum import Enum
|
||||
from pathlib import Path
|
||||
from typing import Any
|
||||
from zoneinfo import ZoneInfo
|
||||
|
||||
from pydantic import BaseModel
|
||||
|
||||
from app.core.config import get_settings
|
||||
|
||||
_SNAKE_PART = re.compile(r"_([a-zA-Z0-9])")
|
||||
_LONG_FIELD_NAMES = {
|
||||
"id",
|
||||
"userId",
|
||||
"creator",
|
||||
"updater",
|
||||
"createUserId",
|
||||
"updateUserId",
|
||||
"createDateTimestamp",
|
||||
"createTime",
|
||||
"createTimeFrom",
|
||||
"createTimeTo",
|
||||
"fileSize",
|
||||
"lastConnectedAtTimestamp",
|
||||
"pid",
|
||||
"reportTime",
|
||||
"size",
|
||||
"timestamp",
|
||||
"tokenCount",
|
||||
"tokenNum",
|
||||
"totalDocCount",
|
||||
"totalTokenCount",
|
||||
"updateTime",
|
||||
}
|
||||
|
||||
|
||||
class JavaMap(dict[str, Any]):
|
||||
"""Marker for Java ``Map`` payloads whose keys Jackson leaves untouched."""
|
||||
|
||||
|
||||
def preserve_java_map_keys(value: Any) -> Any:
|
||||
"""Recursively mark a dynamic Java Map/List graph as key-preserving."""
|
||||
|
||||
if isinstance(value, Mapping):
|
||||
return JavaMap({str(key): preserve_java_map_keys(item) for key, item in value.items()})
|
||||
if isinstance(value, Sequence) and not isinstance(value, str | bytes | bytearray):
|
||||
return [preserve_java_map_keys(item) for item in value]
|
||||
return value
|
||||
|
||||
|
||||
def snake_to_camel(value: str) -> str:
|
||||
return _SNAKE_PART.sub(lambda match: match.group(1).upper(), value)
|
||||
|
||||
|
||||
def _is_long_field(name: str | None) -> bool:
|
||||
if not name:
|
||||
return False
|
||||
return name in _LONG_FIELD_NAMES or name.endswith("Id") or name.endswith("Ids")
|
||||
|
||||
|
||||
def java_compatible(value: Any, *, field_name: str | None = None) -> Any:
|
||||
if value is None or isinstance(value, str | bool | float):
|
||||
return value
|
||||
if isinstance(value, BaseModel):
|
||||
return java_compatible(value.model_dump(by_alias=True, exclude_unset=False), field_name=field_name)
|
||||
if dataclasses.is_dataclass(value) and not isinstance(value, type):
|
||||
return java_compatible(dataclasses.asdict(value), field_name=field_name)
|
||||
if isinstance(value, Enum):
|
||||
return java_compatible(value.value, field_name=field_name)
|
||||
if isinstance(value, datetime):
|
||||
timezone = ZoneInfo(get_settings().timezone)
|
||||
localized = value.astimezone(timezone) if value.tzinfo else value
|
||||
return localized.strftime("%Y-%m-%d %H:%M:%S")
|
||||
if isinstance(value, date):
|
||||
return value.strftime("%Y-%m-%d")
|
||||
if isinstance(value, time):
|
||||
return value.strftime("%H:%M:%S")
|
||||
if isinstance(value, Decimal):
|
||||
return float(value)
|
||||
if isinstance(value, int):
|
||||
return str(value) if _is_long_field(field_name) or not -(2**31) <= value < 2**31 else value
|
||||
if isinstance(value, bytes):
|
||||
return value
|
||||
if isinstance(value, Path):
|
||||
return str(value)
|
||||
if isinstance(value, JavaMap):
|
||||
return {
|
||||
str(raw_key): java_compatible(item, field_name=snake_to_camel(str(raw_key)))
|
||||
for raw_key, item in value.items()
|
||||
}
|
||||
if isinstance(value, Mapping):
|
||||
result: dict[str, Any] = {}
|
||||
for raw_key, item in value.items():
|
||||
key = snake_to_camel(str(raw_key))
|
||||
result[key] = java_compatible(item, field_name=key)
|
||||
return result
|
||||
if isinstance(value, Sequence) and not isinstance(value, str | bytes | bytearray):
|
||||
return [java_compatible(item, field_name=field_name) for item in value]
|
||||
return value
|
||||
@@ -0,0 +1 @@
|
||||
"""Outbound integrations used by the FastAPI manager service."""
|
||||
@@ -0,0 +1,68 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Any
|
||||
|
||||
import httpx
|
||||
|
||||
SUMMARY_PROMPT = """你是一个经验丰富的记忆总结者,擅长将对话内容进行总结摘要,遵循以下规则:
|
||||
1、总结用户的重要信息,以便在未来的对话中提供更个性化的服务
|
||||
2、不要重复总结,不要遗忘之前记忆,除非原来的记忆超过了1800字,否则不要遗忘、不要压缩用户的历史记忆
|
||||
3、用户操控的设备音量、播放音乐、天气、退出、不想对话等和用户本身无关的内容,这些信息不需要加入到总结中
|
||||
4、聊天内容中的今天的日期时间、今天的天气情况与用户事件无关的数据,这些信息如果当成记忆存储会影响后续对话,这些信息不需要加入到总结中
|
||||
5、不要把设备操控的成果结果和失败结果加入到总结中,也不要把用户的一些废话加入到总结中
|
||||
6、不要为了总结而总结,如果用户的聊天没有意义,请返回原来的历史记录也是可以的
|
||||
7、只需要返回总结摘要,严格控制在1800字内
|
||||
8、不要包含代码、xml,不需要解释、注释和说明,保存记忆时仅从对话提取信息,不要混入示例内容
|
||||
9、如果提供了历史记忆,请将新对话内容与历史记忆进行智能合并,保留有价值的历史信息,同时添加新的重要信息
|
||||
|
||||
历史记忆:
|
||||
{history_memory}
|
||||
|
||||
新对话内容:
|
||||
{conversation}"""
|
||||
TITLE_PROMPT = (
|
||||
"请根据以下对话内容,生成一个简洁的会话标题(约15字以内),只返回标题,不要包含任何解释或标点符号:\n{conversation}"
|
||||
)
|
||||
|
||||
|
||||
def _apply_thinking_policy(base_url: str, request: dict[str, Any]) -> None:
|
||||
if "aliyuncs.com" in base_url:
|
||||
request["enable_thinking"] = False
|
||||
elif any(domain in base_url for domain in ("bigmodel.cn", "moonshot.cn", "volces.com")):
|
||||
request["thinking"] = {"type": "disabled"}
|
||||
|
||||
|
||||
async def openai_completion(
|
||||
config: dict[str, Any],
|
||||
prompt: str,
|
||||
*,
|
||||
temperature: float,
|
||||
max_tokens: int,
|
||||
timeout: float,
|
||||
) -> str | None:
|
||||
base_url = str(config.get("base_url") or "")
|
||||
api_key = str(config.get("api_key") or "")
|
||||
if not base_url.strip() or not api_key.strip():
|
||||
return None
|
||||
api_url = base_url if base_url.endswith("/chat/completions") else f"{base_url.rstrip('/')}/chat/completions"
|
||||
request: dict[str, Any] = {
|
||||
"model": config.get("model_name") or "gpt-3.5-turbo",
|
||||
"messages": [{"role": "user", "content": prompt}],
|
||||
"temperature": temperature,
|
||||
"max_tokens": max_tokens,
|
||||
}
|
||||
_apply_thinking_policy(base_url, request)
|
||||
async with httpx.AsyncClient(timeout=timeout) as client:
|
||||
response = await client.post(
|
||||
api_url,
|
||||
json=request,
|
||||
headers={"Content-Type": "application/json", "Authorization": f"Bearer {api_key}"},
|
||||
)
|
||||
response.raise_for_status()
|
||||
payload = response.json()
|
||||
choices = payload.get("choices") if isinstance(payload, dict) else None
|
||||
if not isinstance(choices, list) or not choices:
|
||||
return None
|
||||
message = choices[0].get("message") if isinstance(choices[0], dict) else None
|
||||
content = message.get("content") if isinstance(message, dict) else None
|
||||
return str(content) if content is not None else None
|
||||
@@ -0,0 +1,106 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import base64
|
||||
import hashlib
|
||||
import json
|
||||
from typing import Any
|
||||
from urllib.parse import quote_plus, urlsplit, urlunsplit
|
||||
|
||||
from cryptography.hazmat.primitives import padding
|
||||
from cryptography.hazmat.primitives.ciphers import Cipher, algorithms, modes
|
||||
from websockets.asyncio.client import connect
|
||||
|
||||
|
||||
def _java_aes_key(value: str) -> bytes:
|
||||
raw = value.encode("utf-8")
|
||||
if len(raw) in {16, 24, 32}:
|
||||
return raw
|
||||
return raw[:32].ljust(32, b"\x00")
|
||||
|
||||
|
||||
def encrypt_agent_token(agent_id: str, key: str) -> str:
|
||||
digest = hashlib.md5(agent_id.encode("utf-8"), usedforsecurity=False).hexdigest()
|
||||
plain_text = f'{{"agentId": "{digest}"}}'.encode()
|
||||
padder = padding.PKCS7(128).padder()
|
||||
padded = padder.update(plain_text) + padder.finalize()
|
||||
# ECB is required for byte-for-byte compatibility with Java AES/ECB/PKCS5Padding.
|
||||
encryptor = Cipher(algorithms.AES(_java_aes_key(key)), modes.ECB()).encryptor() # noqa: S305
|
||||
encrypted = encryptor.update(padded) + encryptor.finalize()
|
||||
return base64.b64encode(encrypted).decode("ascii")
|
||||
|
||||
|
||||
def build_agent_mcp_address(endpoint: str | None, agent_id: str) -> str | None:
|
||||
if endpoint is None or not endpoint.strip() or endpoint == "null":
|
||||
return None
|
||||
parsed = urlsplit(endpoint)
|
||||
if not parsed.scheme or not parsed.netloc:
|
||||
raise ValueError("mcp的地址存在错误,请进入参数管理修改mcp接入点地址")
|
||||
marker = "key="
|
||||
marker_index = parsed.query.find(marker)
|
||||
# Java takes everything following the first key= marker, including subsequent query text.
|
||||
key = parsed.query[marker_index + len(marker) :] if marker_index >= 0 else parsed.query[3:]
|
||||
ws_scheme = "wss" if parsed.scheme == "https" else "ws"
|
||||
path = parsed.path
|
||||
parent = path[: path.rfind("/")] if "/" in path else ""
|
||||
base = urlunsplit((ws_scheme, parsed.netloc, parent, "", "")).rstrip("/")
|
||||
token = quote_plus(encrypt_agent_token(agent_id, key), safe="")
|
||||
return f"{base}/mcp/?token={token}"
|
||||
|
||||
|
||||
INITIALIZE_REQUEST = {
|
||||
"jsonrpc": "2.0",
|
||||
"method": "initialize",
|
||||
"params": {
|
||||
"protocolVersion": "2024-11-05",
|
||||
"capabilities": {"roots": {"listChanged": False}, "sampling": {}},
|
||||
"clientInfo": {"name": "xz-mcp-broker", "version": "0.0.1"},
|
||||
},
|
||||
"id": 1,
|
||||
}
|
||||
INITIALIZED_NOTIFICATION = {"jsonrpc": "2.0", "method": "notifications/initialized"}
|
||||
TOOLS_REQUEST = {"jsonrpc": "2.0", "method": "tools/list", "params": None, "id": 2}
|
||||
|
||||
|
||||
async def _receive_matching(websocket: Any, request_id: int, timeout: float) -> dict[str, Any] | None:
|
||||
async def receive() -> dict[str, Any] | None:
|
||||
async for message in websocket:
|
||||
try:
|
||||
value = json.loads(message)
|
||||
except (TypeError, json.JSONDecodeError):
|
||||
continue
|
||||
if isinstance(value, dict) and value.get("id") == request_id:
|
||||
return value
|
||||
return None
|
||||
|
||||
return await asyncio.wait_for(receive(), timeout=timeout)
|
||||
|
||||
|
||||
async def list_mcp_tools(address: str, *, connect_timeout: float = 8.0, session_timeout: float = 10.0) -> list[str]:
|
||||
call_address = address.replace("/mcp/", "/call/")
|
||||
try:
|
||||
async with connect(
|
||||
call_address,
|
||||
open_timeout=connect_timeout,
|
||||
max_size=1024 * 1024,
|
||||
close_timeout=1,
|
||||
) as websocket:
|
||||
await websocket.send(json.dumps(INITIALIZE_REQUEST, ensure_ascii=False, separators=(",", ":")))
|
||||
initialized = await _receive_matching(websocket, 1, session_timeout)
|
||||
if not initialized or "result" not in initialized or "error" in initialized:
|
||||
return []
|
||||
await websocket.send(json.dumps(INITIALIZED_NOTIFICATION, separators=(",", ":")))
|
||||
await websocket.send(json.dumps(TOOLS_REQUEST, separators=(",", ":")))
|
||||
response = await _receive_matching(websocket, 2, session_timeout)
|
||||
if not response or "error" in response:
|
||||
return []
|
||||
result = response.get("result")
|
||||
tools = result.get("tools") if isinstance(result, dict) else None
|
||||
if not isinstance(tools, list):
|
||||
return []
|
||||
return sorted(
|
||||
item["name"] for item in tools if isinstance(item, dict) and isinstance(item.get("name"), str)
|
||||
)
|
||||
# Java treats every connect/protocol/parse failure as an empty tool list.
|
||||
except Exception:
|
||||
return []
|
||||
@@ -0,0 +1,63 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import hashlib
|
||||
import json
|
||||
from datetime import date, datetime, timedelta, timezone
|
||||
from typing import Any
|
||||
|
||||
import httpx
|
||||
|
||||
|
||||
class MqttGatewayError(RuntimeError):
|
||||
def __init__(self, message: str, status_code: int | None = None):
|
||||
super().__init__(message)
|
||||
self.status_code = status_code
|
||||
|
||||
|
||||
def daily_authorization_tokens(signature_key: str, now: datetime | None = None) -> list[str]:
|
||||
if not signature_key.strip() or signature_key.strip().lower() == "null":
|
||||
raise MqttGatewayError("MQTT Gateway signature key is empty")
|
||||
instant = now or datetime.now(tz=timezone.utc)
|
||||
utc_date = instant.astimezone(timezone.utc).date()
|
||||
dates: tuple[date, date, date] = (utc_date, utc_date - timedelta(days=1), utc_date + timedelta(days=1))
|
||||
return [hashlib.sha256(f"{value.isoformat()}{signature_key}".encode()).hexdigest() for value in dates]
|
||||
|
||||
|
||||
async def post_json(
|
||||
url: str,
|
||||
body: Any,
|
||||
signature_key: str,
|
||||
*,
|
||||
timeout_seconds: float,
|
||||
now: datetime | None = None,
|
||||
client: httpx.AsyncClient | None = None,
|
||||
) -> str:
|
||||
encoded = json.dumps(body, ensure_ascii=False, separators=(",", ":")).encode("utf-8")
|
||||
owns_client = client is None
|
||||
selected = client or httpx.AsyncClient()
|
||||
last_unauthorized: int | None = None
|
||||
try:
|
||||
for token in daily_authorization_tokens(signature_key, now):
|
||||
response = await selected.post(
|
||||
url,
|
||||
content=encoded,
|
||||
headers={"Content-Type": "application/json", "Authorization": f"Bearer {token}"},
|
||||
timeout=timeout_seconds,
|
||||
)
|
||||
if response.status_code == 401:
|
||||
last_unauthorized = response.status_code
|
||||
continue
|
||||
if not 200 <= response.status_code < 300:
|
||||
raise MqttGatewayError(
|
||||
f"MQTT Gateway request failed with HTTP status {response.status_code}",
|
||||
response.status_code,
|
||||
)
|
||||
return response.text
|
||||
finally:
|
||||
if owns_client:
|
||||
await selected.aclose()
|
||||
raise MqttGatewayError(
|
||||
"MQTT Gateway rejected all daily authorization tokens"
|
||||
+ ("" if last_unauthorized is None else f" (HTTP {last_unauthorized})"),
|
||||
last_unauthorized,
|
||||
)
|
||||
@@ -0,0 +1,490 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from collections.abc import Mapping
|
||||
from typing import Any
|
||||
|
||||
import httpx
|
||||
from fastapi import UploadFile
|
||||
|
||||
from app.core.errors import AppError
|
||||
|
||||
_DOCUMENT_CHUNK_METHODS = {
|
||||
"naive",
|
||||
"manual",
|
||||
"qa",
|
||||
"table",
|
||||
"paper",
|
||||
"book",
|
||||
"laws",
|
||||
"presentation",
|
||||
"picture",
|
||||
"one",
|
||||
"knowledge_graph",
|
||||
"email",
|
||||
}
|
||||
_RUN_STATUSES = {"UNSTART", "RUNNING", "CANCEL", "DONE", "FAIL"}
|
||||
_DOCUMENT_PARSER_FIELDS = (
|
||||
"chunk_token_num",
|
||||
"delimiter",
|
||||
"layout_recognize",
|
||||
"html4excel",
|
||||
"auto_keywords",
|
||||
"auto_questions",
|
||||
"topn_tags",
|
||||
"raptor",
|
||||
"graphrag",
|
||||
)
|
||||
_DATASET_PARSER_FIELDS = (
|
||||
"chunk_token_num",
|
||||
"delimiter",
|
||||
"layout_recognize",
|
||||
"html4excel",
|
||||
"auto_keywords",
|
||||
"auto_questions",
|
||||
)
|
||||
|
||||
|
||||
class RAGFlowClient:
|
||||
"""Async equivalent of the Java RAGFlow adapter and its wire contract."""
|
||||
|
||||
def __init__(self, config: Mapping[str, Any]):
|
||||
self.config = dict(config)
|
||||
self.base_url = str(config.get("base_url") or config.get("baseUrl") or "").rstrip("/")
|
||||
self.api_key = str(config.get("api_key") or config.get("apiKey") or "")
|
||||
raw_timeout = config.get("timeout")
|
||||
if raw_timeout is None:
|
||||
self.timeout = 30.0
|
||||
else:
|
||||
try:
|
||||
self.timeout = float(int(str(raw_timeout)))
|
||||
except (TypeError, ValueError):
|
||||
self.timeout = 30.0
|
||||
self._validate(config)
|
||||
|
||||
def _validate(self, config: Mapping[str, Any]) -> None:
|
||||
if not config:
|
||||
raise AppError(10164)
|
||||
if not self.base_url.strip():
|
||||
raise AppError(10171)
|
||||
if not self.api_key.strip():
|
||||
raise AppError(10172)
|
||||
if "你" in self.api_key:
|
||||
raise AppError(10173)
|
||||
if not self.base_url.startswith(("http://", "https://")):
|
||||
raise AppError(10174)
|
||||
adapter_type = "ragflow" if "type" not in config else str(config.get("type"))
|
||||
if adapter_type != "ragflow":
|
||||
raise AppError(10184, params=(f"适配器类型未注册: {adapter_type}",))
|
||||
|
||||
async def request(
|
||||
self,
|
||||
method: str,
|
||||
endpoint: str,
|
||||
*,
|
||||
params: Mapping[str, Any] | None = None,
|
||||
json_body: Any = None,
|
||||
files: Mapping[str, Any] | None = None,
|
||||
data: Mapping[str, Any] | None = None,
|
||||
) -> dict[str, Any]:
|
||||
headers = {"Authorization": f"Bearer {self.api_key}"}
|
||||
if files is None:
|
||||
headers["Content-Type"] = "application/json"
|
||||
headers["Accept-Charset"] = "utf-8"
|
||||
normalized_params = {
|
||||
key: self._query_value(value) for key, value in (params or {}).items() if value is not None
|
||||
}
|
||||
try:
|
||||
async with httpx.AsyncClient(timeout=self.timeout) as client:
|
||||
response = await client.request(
|
||||
method,
|
||||
self.base_url + endpoint,
|
||||
params=normalized_params,
|
||||
json=json_body,
|
||||
files=files,
|
||||
data=data,
|
||||
headers=headers,
|
||||
)
|
||||
response.raise_for_status()
|
||||
payload = response.json()
|
||||
except (httpx.HTTPError, ValueError) as exc:
|
||||
raise AppError(10167, params=(f"Request Failed: {exc}",)) from exc
|
||||
if not isinstance(payload, dict):
|
||||
raise AppError(10167, params=("Invalid Response",))
|
||||
code = payload.get("code")
|
||||
if code is not None:
|
||||
if isinstance(code, bool) or not isinstance(code, int):
|
||||
raise AppError(10167, params=("Request Failed: invalid response code type",))
|
||||
if code != 0:
|
||||
message = payload.get("message")
|
||||
if message is not None and not isinstance(message, str):
|
||||
raise AppError(10167, params=("Request Failed: invalid response message type",))
|
||||
raise AppError(10167, params=(message or "Unknown RAGFlow Error",))
|
||||
return dict(payload)
|
||||
|
||||
@staticmethod
|
||||
def _query_value(value: Any) -> Any:
|
||||
if isinstance(value, bool):
|
||||
return str(value).lower()
|
||||
if isinstance(value, list):
|
||||
# Java List.toString() is what the baseline URL builder sends.
|
||||
return "[" + ", ".join(str(item) for item in value) + "]"
|
||||
return value
|
||||
|
||||
async def dataset_info(self, dataset_id: str) -> dict[str, Any] | None:
|
||||
payload = await self.request(
|
||||
"GET", "/api/v1/datasets", params={"id": dataset_id, "page": 1, "page_size": 1}
|
||||
)
|
||||
data = payload.get("data")
|
||||
if isinstance(data, list) and data and isinstance(data[0], dict):
|
||||
return _normalize_dataset_info(data[0])
|
||||
return None
|
||||
|
||||
async def create_dataset(self, body: dict[str, Any]) -> dict[str, Any]:
|
||||
body = dict(body)
|
||||
body["permission"] = "me" if _blank(body.get("permission")) else body.get("permission")
|
||||
body["chunk_method"] = "naive" if _blank(body.get("chunk_method")) else body.get("chunk_method")
|
||||
if _blank(body.get("embedding_model")):
|
||||
configured_model = self.config.get("embedding_model", self.config.get("embeddingModel"))
|
||||
body["embedding_model"] = None if _blank(configured_model) else configured_model
|
||||
body["avatar"] = body.get("avatar") if not _blank(body.get("avatar")) else (
|
||||
"data:image/png;base64,iVBORw0KGgoAAAANSUhEUgAAAAEAAAABCAYAAAAfFcSJAAAADUlEQVR42mP8z8BQDwAEhQGAhKmMIQAAAABJRU5ErkJggg=="
|
||||
)
|
||||
body["parser_config"] = _normalize_parser_config(
|
||||
body.get("parser_config"), fields=_DATASET_PARSER_FIELDS
|
||||
)
|
||||
payload = await self.request("POST", "/api/v1/datasets", json_body=body)
|
||||
data = payload.get("data")
|
||||
if not isinstance(data, dict) or not data.get("id"):
|
||||
raise AppError(10167, params=("Invalid response from createDataset: missing data object",))
|
||||
return _normalize_dataset_info(data)
|
||||
|
||||
async def update_dataset(self, dataset_id: str, body: dict[str, Any]) -> dict[str, Any] | None:
|
||||
body = dict(body)
|
||||
body["parser_config"] = _normalize_parser_config(
|
||||
body.get("parser_config"), fields=_DATASET_PARSER_FIELDS
|
||||
)
|
||||
payload = await self.request("PUT", f"/api/v1/datasets/{dataset_id}", json_body=body)
|
||||
return _normalize_dataset_info(payload["data"]) if isinstance(payload.get("data"), dict) else None
|
||||
|
||||
async def delete_datasets(self, ids: list[str]) -> Any:
|
||||
return (await self.request("DELETE", "/api/v1/datasets", json_body={"ids": ids})).get("data")
|
||||
|
||||
async def documents(
|
||||
self,
|
||||
dataset_id: str,
|
||||
*,
|
||||
page: int = 1,
|
||||
page_size: int = 10,
|
||||
name: str | None = None,
|
||||
status: str | None = None,
|
||||
document_id: str | None = None,
|
||||
) -> tuple[list[dict[str, Any]], int]:
|
||||
params: dict[str, Any] = {"page": page, "page_size": page_size}
|
||||
if name:
|
||||
params["name"] = name
|
||||
if status:
|
||||
status_number = int(status) if status.lstrip("-").isdigit() else None
|
||||
names = {0: "UNSTART", 1: "RUNNING", 2: "CANCEL", 3: "DONE", 4: "FAIL"}
|
||||
params["run"] = [names[status_number]] if status_number in names else []
|
||||
if document_id:
|
||||
params["id"] = document_id
|
||||
payload = await self.request("GET", f"/api/v1/datasets/{dataset_id}/documents", params=params)
|
||||
data = payload.get("data")
|
||||
if not isinstance(data, dict):
|
||||
return [], 0
|
||||
docs = data.get("docs")
|
||||
rows: list[dict[str, Any]] = []
|
||||
if isinstance(docs, list):
|
||||
for item in docs:
|
||||
if not isinstance(item, dict):
|
||||
continue
|
||||
try:
|
||||
rows.append(_normalize_upload_document(item))
|
||||
except AppError:
|
||||
# The Java adapter skips an individual document whose
|
||||
# strong DTO conversion fails and keeps the rest of page.
|
||||
continue
|
||||
return rows, int(data.get("total") or 0)
|
||||
|
||||
async def upload_document(
|
||||
self,
|
||||
dataset_id: str,
|
||||
file: UploadFile,
|
||||
content: bytes,
|
||||
*,
|
||||
name: str,
|
||||
meta_fields: dict[str, Any] | None,
|
||||
chunk_method: str | None,
|
||||
parser_config: dict[str, Any] | None,
|
||||
) -> dict[str, Any]:
|
||||
import json
|
||||
|
||||
form: dict[str, Any] = {"name": name}
|
||||
if meta_fields:
|
||||
form["meta"] = json.dumps(meta_fields, ensure_ascii=False, separators=(",", ":"))
|
||||
if not _blank(chunk_method):
|
||||
normalized_method = str(chunk_method).lower()
|
||||
if normalized_method in _DOCUMENT_CHUNK_METHODS:
|
||||
form["chunk_method"] = normalized_method
|
||||
normalized_parser = (
|
||||
_normalize_parser_config(
|
||||
parser_config,
|
||||
fields=_DOCUMENT_PARSER_FIELDS,
|
||||
validate_layout=True,
|
||||
)
|
||||
if parser_config
|
||||
else None
|
||||
)
|
||||
if normalized_parser is not None:
|
||||
form["parser_config"] = json.dumps(normalized_parser, ensure_ascii=False, separators=(",", ":"))
|
||||
payload = await self.request(
|
||||
"POST",
|
||||
f"/api/v1/datasets/{dataset_id}/documents",
|
||||
files={"file": (file.filename or name, content, file.content_type or "application/octet-stream")},
|
||||
data=form,
|
||||
)
|
||||
data = payload.get("data")
|
||||
if isinstance(data, list) and data and isinstance(data[0], dict):
|
||||
return _normalize_upload_document(data[0])
|
||||
if isinstance(data, dict):
|
||||
return _normalize_upload_document(data)
|
||||
raise AppError(10167, params=("远程上传成功但未返回有效 DocumentID",))
|
||||
|
||||
async def delete_documents(self, dataset_id: str, ids: list[str]) -> None:
|
||||
await self.request(
|
||||
"DELETE", f"/api/v1/datasets/{dataset_id}/documents", json_body={"ids": ids}
|
||||
)
|
||||
|
||||
async def parse_documents(self, dataset_id: str, document_ids: list[str]) -> None:
|
||||
await self.request(
|
||||
"POST",
|
||||
f"/api/v1/datasets/{dataset_id}/chunks",
|
||||
json_body={"document_ids": document_ids},
|
||||
)
|
||||
|
||||
async def chunks(
|
||||
self, dataset_id: str, document_id: str, params: Mapping[str, Any]
|
||||
) -> dict[str, Any]:
|
||||
payload = await self.request(
|
||||
"GET", f"/api/v1/datasets/{dataset_id}/documents/{document_id}/chunks", params=params
|
||||
)
|
||||
data = payload.get("data")
|
||||
if not isinstance(data, dict):
|
||||
return {"chunks": [], "doc": None, "total": 0}
|
||||
try:
|
||||
return _normalize_chunk_list(data)
|
||||
except (TypeError, ValueError) as exc:
|
||||
raise AppError(10167, params=(str(exc),)) from exc
|
||||
|
||||
async def retrieval(self, body: dict[str, Any]) -> dict[str, Any]:
|
||||
payload = await self.request("POST", "/api/v1/retrieval", json_body=body)
|
||||
data = payload.get("data")
|
||||
if not isinstance(data, dict):
|
||||
return {"chunks": [], "doc_aggs": [], "total": 0}
|
||||
try:
|
||||
return _normalize_retrieval_result(data)
|
||||
except (TypeError, ValueError) as exc:
|
||||
raise AppError(10167, params=(str(exc),)) from exc
|
||||
|
||||
|
||||
def _blank(value: Any) -> bool:
|
||||
return value is None or (isinstance(value, str) and not value.strip())
|
||||
|
||||
|
||||
def _normalize_parser_config(
|
||||
value: Any,
|
||||
*,
|
||||
fields: tuple[str, ...],
|
||||
validate_layout: bool = False,
|
||||
) -> dict[str, Any] | None:
|
||||
if value is None:
|
||||
return None
|
||||
if not isinstance(value, Mapping):
|
||||
raise ValueError("parser_config must be an object")
|
||||
result = {key: value.get(key) for key in fields}
|
||||
if validate_layout and result.get("layout_recognize") not in {None, "DeepDOC", "Simple"}:
|
||||
raise ValueError("invalid layout_recognize")
|
||||
if "raptor" in result and result["raptor"] is not None:
|
||||
nested = result["raptor"]
|
||||
if not isinstance(nested, Mapping):
|
||||
raise ValueError("raptor must be an object")
|
||||
result["raptor"] = {"use_raptor": nested.get("use_raptor")}
|
||||
if "graphrag" in result and result["graphrag"] is not None:
|
||||
nested = result["graphrag"]
|
||||
if not isinstance(nested, Mapping):
|
||||
raise ValueError("graphrag must be an object")
|
||||
result["graphrag"] = {"use_graphrag": nested.get("use_graphrag")}
|
||||
return result
|
||||
|
||||
|
||||
def _normalize_upload_document(value: Mapping[str, Any]) -> dict[str, Any]:
|
||||
result = dict(value)
|
||||
try:
|
||||
if result.get("parser_config") is not None:
|
||||
result["parser_config"] = _normalize_parser_config(
|
||||
result["parser_config"], fields=_DOCUMENT_PARSER_FIELDS, validate_layout=True
|
||||
)
|
||||
except (TypeError, ValueError) as exc:
|
||||
raise AppError(10167, params=("远程上传成功但未返回有效 DocumentID",)) from exc
|
||||
chunk_method = result.get("chunk_method")
|
||||
if chunk_method is not None:
|
||||
normalized_method = str(chunk_method).lower()
|
||||
if normalized_method not in _DOCUMENT_CHUNK_METHODS:
|
||||
raise AppError(10167, params=("远程上传成功但未返回有效 DocumentID",))
|
||||
result["chunk_method"] = normalized_method
|
||||
run = result.get("run")
|
||||
if run is not None and str(run) not in _RUN_STATUSES:
|
||||
raise AppError(10167, params=("远程上传成功但未返回有效 DocumentID",))
|
||||
return result
|
||||
|
||||
|
||||
def _normalize_dataset_info(value: Mapping[str, Any]) -> dict[str, Any]:
|
||||
"""Apply Jackson's DatasetDTO.InfoVO unknown-field and type boundary."""
|
||||
fields = (
|
||||
"id",
|
||||
"name",
|
||||
"avatar",
|
||||
"tenant_id",
|
||||
"description",
|
||||
"embedding_model",
|
||||
"permission",
|
||||
"chunk_method",
|
||||
"parser_config",
|
||||
"chunk_count",
|
||||
"document_count",
|
||||
"create_time",
|
||||
"update_time",
|
||||
"token_num",
|
||||
"create_date",
|
||||
"update_date",
|
||||
)
|
||||
try:
|
||||
result = {field: value.get(field) for field in fields}
|
||||
result["parser_config"] = _normalize_parser_config(
|
||||
result.get("parser_config"), fields=_DATASET_PARSER_FIELDS
|
||||
)
|
||||
for field in ("chunk_count", "document_count", "create_time", "update_time", "token_num"):
|
||||
raw_value = result[field]
|
||||
if raw_value is not None:
|
||||
result[field] = int(raw_value)
|
||||
return result
|
||||
except (TypeError, ValueError) as exc:
|
||||
raise AppError(10167, params=(str(exc),)) from exc
|
||||
|
||||
|
||||
def _nullable_object(value: Any, fields: tuple[str, ...]) -> dict[str, Any] | None:
|
||||
if value is None:
|
||||
return None
|
||||
if not isinstance(value, Mapping):
|
||||
raise TypeError("response object has an invalid shape")
|
||||
return {field: value.get(field) for field in fields}
|
||||
|
||||
|
||||
def _normalize_chunk_list(data: Mapping[str, Any]) -> dict[str, Any]:
|
||||
chunk_fields = (
|
||||
"id",
|
||||
"content",
|
||||
"document_id",
|
||||
"docnm_kwd",
|
||||
"important_keywords",
|
||||
"questions",
|
||||
"image_id",
|
||||
"dataset_id",
|
||||
"available",
|
||||
"positions",
|
||||
"token",
|
||||
)
|
||||
raw_chunks = data.get("chunks")
|
||||
if raw_chunks is None:
|
||||
chunks: list[dict[str, Any]] = []
|
||||
elif isinstance(raw_chunks, list):
|
||||
chunks = []
|
||||
for item in raw_chunks:
|
||||
normalized = _nullable_object(item, chunk_fields)
|
||||
if normalized is not None:
|
||||
chunks.append(normalized)
|
||||
else:
|
||||
raise TypeError("chunks must be an array")
|
||||
|
||||
doc_fields = (
|
||||
"id",
|
||||
"thumbnail",
|
||||
"dataset_id",
|
||||
"chunk_method",
|
||||
"pipeline_id",
|
||||
"parser_config",
|
||||
"source_type",
|
||||
"type",
|
||||
"created_by",
|
||||
"name",
|
||||
"location",
|
||||
"size",
|
||||
"token_count",
|
||||
"chunk_count",
|
||||
"progress",
|
||||
"progress_msg",
|
||||
"process_begin_at",
|
||||
"process_duration",
|
||||
"meta_fields",
|
||||
"suffix",
|
||||
"run",
|
||||
"status",
|
||||
"create_time",
|
||||
"create_date",
|
||||
"update_time",
|
||||
"update_date",
|
||||
)
|
||||
doc = _nullable_object(data.get("doc"), doc_fields)
|
||||
if doc is not None:
|
||||
doc["parser_config"] = _normalize_parser_config(
|
||||
doc.get("parser_config"), fields=_DOCUMENT_PARSER_FIELDS, validate_layout=True
|
||||
)
|
||||
if doc.get("chunk_count") is not None:
|
||||
# DocumentDTO.InfoVO.chunkCount is Long, unlike the Integer field
|
||||
# on KnowledgeFilesDTO used by the document-list endpoint.
|
||||
doc["chunk_count"] = str(doc["chunk_count"])
|
||||
if doc.get("run") is not None and str(doc["run"]) not in _RUN_STATUSES:
|
||||
raise ValueError("invalid document run status")
|
||||
# ChunkDTO.ListVO.total is Long and therefore uses the Java global Long
|
||||
# serializer even for small values (including the adapter's default 0L).
|
||||
return {"chunks": chunks, "doc": doc, "total": str(int(data.get("total") or 0))}
|
||||
|
||||
|
||||
def _normalize_retrieval_result(data: Mapping[str, Any]) -> dict[str, Any]:
|
||||
hit_fields = (
|
||||
"id",
|
||||
"content",
|
||||
"document_id",
|
||||
"dataset_id",
|
||||
"document_name",
|
||||
"document_keyword",
|
||||
"similarity",
|
||||
"vector_similarity",
|
||||
"term_similarity",
|
||||
"index",
|
||||
"highlight",
|
||||
"important_keywords",
|
||||
"questions",
|
||||
"image_id",
|
||||
"positions",
|
||||
)
|
||||
agg_fields = ("doc_name", "doc_id", "count")
|
||||
|
||||
def normalize_list(raw: Any, fields: tuple[str, ...], name: str) -> list[dict[str, Any]]:
|
||||
if raw is None:
|
||||
return []
|
||||
if not isinstance(raw, list):
|
||||
raise TypeError(f"{name} must be an array")
|
||||
values: list[dict[str, Any]] = []
|
||||
for item in raw:
|
||||
normalized = _nullable_object(item, fields)
|
||||
if normalized is not None:
|
||||
values.append(normalized)
|
||||
return values
|
||||
|
||||
return {
|
||||
"chunks": normalize_list(data.get("chunks"), hit_fields, "chunks"),
|
||||
"doc_aggs": normalize_list(data.get("doc_aggs"), agg_fields, "doc_aggs"),
|
||||
# RetrievalDTO.ResultVO.total is also Long.
|
||||
"total": str(int(data.get("total") or 0)),
|
||||
}
|
||||
@@ -0,0 +1,96 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import base64
|
||||
import json
|
||||
from dataclasses import dataclass
|
||||
from typing import Any
|
||||
|
||||
import httpx
|
||||
|
||||
|
||||
@dataclass(slots=True)
|
||||
class VoiceCloneProviderError(Exception):
|
||||
code: int
|
||||
message: str
|
||||
|
||||
def __str__(self) -> str:
|
||||
return self.message
|
||||
|
||||
|
||||
class VoiceCloneIntegration:
|
||||
ENDPOINT = "https://openspeech.bytedance.com/api/v1/mega_tts/audio/upload"
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
*,
|
||||
timeout_seconds: float,
|
||||
client: httpx.AsyncClient | None = None,
|
||||
endpoint: str | None = None,
|
||||
):
|
||||
self.timeout_seconds = timeout_seconds
|
||||
self.client = client
|
||||
self.endpoint = endpoint or self.ENDPOINT
|
||||
|
||||
async def train_huoshan(
|
||||
self,
|
||||
*,
|
||||
appid: str,
|
||||
access_token: str,
|
||||
voice: bytes,
|
||||
speaker_id: str,
|
||||
) -> str:
|
||||
request_body: dict[str, Any] = {
|
||||
"appid": appid,
|
||||
"audios": [
|
||||
{
|
||||
"audio_bytes": base64.b64encode(voice).decode("ascii"),
|
||||
"audio_format": "wav",
|
||||
}
|
||||
],
|
||||
"source": 2,
|
||||
"language": 0,
|
||||
"model_type": 1,
|
||||
"speaker_id": speaker_id,
|
||||
}
|
||||
owns_client = self.client is None
|
||||
client = self.client or httpx.AsyncClient()
|
||||
try:
|
||||
response = await client.post(
|
||||
self.endpoint,
|
||||
content=json.dumps(request_body, ensure_ascii=False, separators=(",", ":")).encode("utf-8"),
|
||||
headers={
|
||||
"Content-Type": "application/json",
|
||||
"Authorization": f"Bearer;{access_token}",
|
||||
"Resource-Id": "seed-icl-1.0",
|
||||
},
|
||||
timeout=self.timeout_seconds,
|
||||
)
|
||||
try:
|
||||
payload = response.json()
|
||||
except (json.JSONDecodeError, ValueError) as exc:
|
||||
raise VoiceCloneProviderError(10157, str(exc)) from exc
|
||||
except httpx.HTTPError as exc:
|
||||
raise VoiceCloneProviderError(10157, str(exc)) from exc
|
||||
finally:
|
||||
if owns_client:
|
||||
await client.aclose()
|
||||
|
||||
if not isinstance(payload, dict):
|
||||
raise VoiceCloneProviderError(10156, "响应格式错误,缺少BaseResp字段")
|
||||
base_response = payload.get("BaseResp")
|
||||
if isinstance(base_response, dict):
|
||||
raw_status = base_response.get("StatusCode")
|
||||
try:
|
||||
status_code = int(raw_status) if raw_status is not None else None
|
||||
except (TypeError, ValueError):
|
||||
status_code = None
|
||||
returned_speaker = payload.get("speaker_id")
|
||||
if status_code == 0 and isinstance(returned_speaker, str) and returned_speaker.strip():
|
||||
return returned_speaker
|
||||
status_message = base_response.get("StatusMessage")
|
||||
message = str(status_message) if status_message not in (None, "") else "训练失败"
|
||||
raise VoiceCloneProviderError(500, message)
|
||||
payload_message = payload.get("message")
|
||||
if payload_message not in (None, ""):
|
||||
raise VoiceCloneProviderError(500, str(payload_message))
|
||||
raise VoiceCloneProviderError(10156, "响应格式错误,缺少BaseResp字段")
|
||||
@@ -0,0 +1,102 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from collections.abc import Mapping
|
||||
from dataclasses import dataclass
|
||||
from urllib.parse import urlsplit
|
||||
|
||||
import httpx
|
||||
|
||||
|
||||
class VoicePrintIntegrationError(RuntimeError):
|
||||
def __init__(self, code: int, message: str | None = None, params: tuple[object, ...] = ()):
|
||||
super().__init__(message or str(code))
|
||||
self.code = code
|
||||
self.message = message
|
||||
self.params = params
|
||||
|
||||
|
||||
@dataclass(slots=True, frozen=True)
|
||||
class VoicePrintEndpoint:
|
||||
base_url: str
|
||||
authorization: str
|
||||
|
||||
@classmethod
|
||||
def parse(cls, configured_url: str | None) -> VoicePrintEndpoint:
|
||||
if configured_url is None:
|
||||
raise VoicePrintIntegrationError(10084)
|
||||
parsed = urlsplit(configured_url)
|
||||
if not parsed.scheme or not parsed.hostname:
|
||||
raise VoicePrintIntegrationError(10084)
|
||||
marker = "key="
|
||||
marker_index = parsed.query.find(marker)
|
||||
key = parsed.query[marker_index + len(marker) :] if marker_index >= 0 else parsed.query[3:]
|
||||
port = f":{parsed.port}" if parsed.port is not None else ""
|
||||
return cls(f"{parsed.scheme}://{parsed.hostname}{port}", f"Bearer {key}")
|
||||
|
||||
|
||||
class VoicePrintClient:
|
||||
def __init__(self, configured_url: str, *, timeout: float = 10.0, client: httpx.AsyncClient | None = None):
|
||||
self.endpoint = VoicePrintEndpoint.parse(configured_url)
|
||||
self.timeout = timeout
|
||||
self._client = client
|
||||
|
||||
async def _request(
|
||||
self,
|
||||
method: str,
|
||||
path: str,
|
||||
*,
|
||||
data: Mapping[str, str] | None = None,
|
||||
files: Mapping[str, tuple[str, bytes, str]] | None = None,
|
||||
) -> httpx.Response:
|
||||
headers = {"Authorization": self.endpoint.authorization}
|
||||
if self._client is not None:
|
||||
return await self._client.request(
|
||||
method, f"{self.endpoint.base_url}{path}", headers=headers, data=data, files=files
|
||||
)
|
||||
async with httpx.AsyncClient(timeout=self.timeout) as client:
|
||||
return await client.request(
|
||||
method, f"{self.endpoint.base_url}{path}", headers=headers, data=data, files=files
|
||||
)
|
||||
|
||||
async def identify(self, speaker_ids: list[str], audio: bytes) -> tuple[str | None, float | None] | None:
|
||||
if not speaker_ids:
|
||||
return None
|
||||
response = await self._request(
|
||||
"POST",
|
||||
"/voiceprint/identify",
|
||||
data={"speaker_ids": ",".join(speaker_ids)},
|
||||
files={"file": ("VoicePrint.WAV", audio, "application/octet-stream")},
|
||||
)
|
||||
if response.status_code != 200:
|
||||
raise VoicePrintIntegrationError(10091)
|
||||
try:
|
||||
payload = response.json()
|
||||
except ValueError as exc:
|
||||
raise VoicePrintIntegrationError(10091) from exc
|
||||
if not isinstance(payload, dict):
|
||||
return None
|
||||
speaker_id = payload.get("speaker_id")
|
||||
score = payload.get("score")
|
||||
return (
|
||||
str(speaker_id) if speaker_id is not None else None,
|
||||
float(score) if isinstance(score, int | float) else None,
|
||||
)
|
||||
|
||||
async def register(self, speaker_id: str, audio: bytes) -> None:
|
||||
response = await self._request(
|
||||
"POST",
|
||||
"/voiceprint/register",
|
||||
data={"speaker_id": speaker_id},
|
||||
files={"file": ("VoicePrint.WAV", audio, "application/octet-stream")},
|
||||
)
|
||||
if response.status_code != 200:
|
||||
raise VoicePrintIntegrationError(10087)
|
||||
if "true" not in response.text:
|
||||
raise VoicePrintIntegrationError(10088)
|
||||
|
||||
async def cancel(self, speaker_id: str) -> None:
|
||||
response = await self._request("DELETE", f"/voiceprint/{speaker_id}")
|
||||
if response.status_code != 200:
|
||||
raise VoicePrintIntegrationError(10089)
|
||||
if "true" not in response.text:
|
||||
raise VoicePrintIntegrationError(10090)
|
||||
@@ -0,0 +1 @@
|
||||
"""Single-instance background jobs for the manager API."""
|
||||
@@ -0,0 +1,18 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from app.core.config import get_settings
|
||||
from app.core.database import get_session_factory
|
||||
from app.core.redis import distributed_lock
|
||||
from app.repositories.knowledge import KnowledgeRepository
|
||||
from app.services.knowledge import KnowledgeDocumentService
|
||||
|
||||
|
||||
async def sync_running_knowledge_documents() -> int:
|
||||
"""Run one document-status pass under a cross-process Redis lock."""
|
||||
|
||||
settings = get_settings()
|
||||
async with distributed_lock("jobs:knowledge-document-status", settings.job_lock_ttl_seconds) as acquired:
|
||||
if not acquired:
|
||||
return 0
|
||||
async with get_session_factory()() as session:
|
||||
return await KnowledgeDocumentService(KnowledgeRepository(session)).sync_running()
|
||||
@@ -0,0 +1,135 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import logging
|
||||
import os
|
||||
import signal
|
||||
import time
|
||||
from collections.abc import Awaitable, Callable
|
||||
from contextlib import suppress
|
||||
|
||||
from app.core.config import get_settings
|
||||
from app.core.database import configure_database, database_ping, dispose_database
|
||||
from app.core.redis import close_redis, redis_ping
|
||||
from app.jobs.tasks import sync_running_knowledge_documents
|
||||
from app.services.agent import redact_legacy_agent_snapshots
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
async def _wait_or_stop(stop: asyncio.Event, seconds: float) -> bool:
|
||||
try:
|
||||
await asyncio.wait_for(stop.wait(), timeout=seconds)
|
||||
except asyncio.TimeoutError:
|
||||
return False
|
||||
return True
|
||||
|
||||
|
||||
async def _fixed_delay_loop(
|
||||
stop: asyncio.Event,
|
||||
operation: Callable[[], Awaitable[int]],
|
||||
*,
|
||||
name: str,
|
||||
initial_delay: float,
|
||||
delay: float,
|
||||
) -> None:
|
||||
if initial_delay and await _wait_or_stop(stop, initial_delay):
|
||||
return
|
||||
while not stop.is_set():
|
||||
started = time.monotonic()
|
||||
try:
|
||||
changed = await operation()
|
||||
logger.info(
|
||||
"Background job %s completed changed=%s duration_ms=%d",
|
||||
name,
|
||||
changed,
|
||||
(time.monotonic() - started) * 1000,
|
||||
)
|
||||
except Exception:
|
||||
logger.exception("Background job %s failed; the next fixed-delay pass will retry", name)
|
||||
if await _wait_or_stop(stop, delay):
|
||||
return
|
||||
|
||||
|
||||
async def _finish_tasks(tasks: list[asyncio.Task[None]], timeout: float) -> bool:
|
||||
"""Let active jobs finish, then cancel only those exceeding the shutdown budget."""
|
||||
|
||||
_, pending = await asyncio.wait(tasks, timeout=max(timeout, 0.0))
|
||||
if not pending:
|
||||
return True
|
||||
logger.warning(
|
||||
"Graceful job shutdown timed out after %.1f seconds; cancelling %d task(s)",
|
||||
timeout,
|
||||
len(pending),
|
||||
)
|
||||
for task in pending:
|
||||
task.cancel()
|
||||
for task in pending:
|
||||
with suppress(asyncio.CancelledError):
|
||||
await task
|
||||
return False
|
||||
|
||||
|
||||
async def run_worker(stop: asyncio.Event | None = None) -> None:
|
||||
settings = get_settings()
|
||||
os.environ["TZ"] = settings.timezone
|
||||
if hasattr(time, "tzset"):
|
||||
time.tzset()
|
||||
logging.basicConfig(level=settings.log_level, format="%(asctime)s %(levelname)s %(name)s %(message)s")
|
||||
configure_database(settings)
|
||||
selected_stop = stop or asyncio.Event()
|
||||
|
||||
if not settings.allow_start_without_dependencies:
|
||||
if not await database_ping():
|
||||
raise RuntimeError("database readiness check failed")
|
||||
if not await redis_ping():
|
||||
raise RuntimeError("Redis readiness check failed")
|
||||
|
||||
# The retained Java service performs a blocking startup redaction pass,
|
||||
# then compensates rolling-deployment writes after 5 seconds and every
|
||||
# 15 seconds. The standalone worker preserves those timings while the
|
||||
# Redis lock makes multiple worker replicas safe.
|
||||
await redact_legacy_agent_snapshots()
|
||||
tasks = [
|
||||
asyncio.create_task(
|
||||
_fixed_delay_loop(
|
||||
selected_stop,
|
||||
redact_legacy_agent_snapshots,
|
||||
name="agent-snapshot-redaction",
|
||||
initial_delay=5,
|
||||
delay=15,
|
||||
)
|
||||
),
|
||||
asyncio.create_task(
|
||||
_fixed_delay_loop(
|
||||
selected_stop,
|
||||
sync_running_knowledge_documents,
|
||||
name="knowledge-document-status",
|
||||
initial_delay=0,
|
||||
delay=30,
|
||||
)
|
||||
),
|
||||
]
|
||||
try:
|
||||
await selected_stop.wait()
|
||||
finally:
|
||||
await _finish_tasks(tasks, settings.graceful_shutdown_seconds)
|
||||
await close_redis()
|
||||
await dispose_database()
|
||||
|
||||
|
||||
def main() -> None:
|
||||
stop = asyncio.Event()
|
||||
loop = asyncio.new_event_loop()
|
||||
asyncio.set_event_loop(loop)
|
||||
for signal_name in (signal.SIGINT, signal.SIGTERM):
|
||||
with suppress(NotImplementedError):
|
||||
loop.add_signal_handler(signal_name, stop.set)
|
||||
try:
|
||||
loop.run_until_complete(run_worker(stop))
|
||||
finally:
|
||||
loop.close()
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,217 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
import os
|
||||
import time
|
||||
from collections.abc import AsyncIterator, Mapping, Sequence
|
||||
from contextlib import asynccontextmanager
|
||||
from typing import Any
|
||||
|
||||
from fastapi import FastAPI, Request
|
||||
from fastapi.exceptions import RequestValidationError
|
||||
from sqlalchemy.exc import IntegrityError
|
||||
from starlette.exceptions import HTTPException as StarletteHTTPException
|
||||
from starlette.middleware.cors import CORSMiddleware
|
||||
|
||||
from app.core.config import get_settings
|
||||
from app.core.database import configure_database, database_ping, dispose_database
|
||||
from app.core.errors import AppError, ErrorCode
|
||||
from app.core.i18n import message_for
|
||||
from app.core.redis import close_redis, redis_ping
|
||||
from app.core.responses import JavaJSONResponse, error_response, ok
|
||||
from app.core.security import AuthenticationMiddleware
|
||||
from app.routers import application_routers
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
settings = get_settings()
|
||||
|
||||
_MULTIPART_VALIDATION_PATHS = {
|
||||
"/datasets/{dataset_id}/documents",
|
||||
"/otaMag/upload",
|
||||
"/otaMag/uploadAssetsBin",
|
||||
"/voiceClone/upload",
|
||||
}
|
||||
|
||||
|
||||
def _matches_path_template(path: str, template: str) -> bool:
|
||||
path_parts = path.removeprefix(settings.context_path).strip("/").split("/")
|
||||
template_parts = template.strip("/").split("/")
|
||||
return len(path_parts) == len(template_parts) and all(
|
||||
expected.startswith("{") and expected.endswith("}") or actual == expected
|
||||
for actual, expected in zip(path_parts, template_parts, strict=True)
|
||||
)
|
||||
|
||||
|
||||
def _java_required_message(request: Request, errors: Sequence[Mapping[str, Any]]) -> str | None:
|
||||
path = request.url.path.removeprefix(settings.context_path)
|
||||
mappings = (
|
||||
(
|
||||
"/admin/server/emit-action",
|
||||
(("action", "操作不能为空"), ("targetWs", "目标ws地址不能为空")),
|
||||
),
|
||||
("/agent", (("agentName", "智能体名称不能为空"),)),
|
||||
(
|
||||
"/agent/chat-history/report",
|
||||
tuple((field, "不能为空") for field in ("macAddress", "sessionId", "chatType", "content")),
|
||||
),
|
||||
(
|
||||
"/agent/{agentId}/snapshots/{snapshotId}/restore",
|
||||
(("currentStateToken", "不能为空"),),
|
||||
),
|
||||
(
|
||||
"/config/agent-models",
|
||||
(
|
||||
("macAddress", "设备MAC地址不能为空"),
|
||||
("clientId", "客户端ID不能为空"),
|
||||
("selectedModule", "客户端已实例化的模型不能为空"),
|
||||
),
|
||||
),
|
||||
("/config/correct-words", (("macAddress", "设备MAC地址不能为空"),)),
|
||||
(
|
||||
"/device/address-book/alias",
|
||||
(("targetMac", "目标MAC地址不能为空"), ("macAddress", "MAC地址不能为空")),
|
||||
),
|
||||
)
|
||||
missing_fields: set[str] = set()
|
||||
for error in errors:
|
||||
location = tuple(error.get("loc", ()))
|
||||
if error.get("type") == "missing" and location[:1] == ("body",):
|
||||
missing_fields.add(str(location[-1]))
|
||||
if not missing_fields:
|
||||
return None
|
||||
for template, fields in mappings:
|
||||
if _matches_path_template(path, template):
|
||||
return next((message for field, message in fields if field in missing_fields), None)
|
||||
return None
|
||||
|
||||
|
||||
@asynccontextmanager
|
||||
async def lifespan(_: FastAPI) -> AsyncIterator[None]:
|
||||
os.environ["TZ"] = settings.timezone
|
||||
if hasattr(time, "tzset"):
|
||||
time.tzset()
|
||||
settings.upload_dir.mkdir(parents=True, exist_ok=True)
|
||||
configure_database(settings)
|
||||
if not settings.i18n_dir.exists():
|
||||
raise RuntimeError(f"Java i18n resources are missing: {settings.i18n_dir}")
|
||||
if not settings.changelog_path.exists():
|
||||
raise RuntimeError(f"Liquibase source of truth is missing: {settings.changelog_path}")
|
||||
if not settings.allow_start_without_dependencies:
|
||||
if not await database_ping():
|
||||
raise RuntimeError("database readiness check failed")
|
||||
if not await redis_ping():
|
||||
raise RuntimeError("Redis readiness check failed")
|
||||
yield
|
||||
await close_redis()
|
||||
await dispose_database()
|
||||
|
||||
|
||||
app = FastAPI(
|
||||
title="xiaozhi-manager-api",
|
||||
version="0.1.0",
|
||||
docs_url=f"{settings.context_path}/doc.html",
|
||||
openapi_url=f"{settings.context_path}/v3/api-docs",
|
||||
redoc_url=None,
|
||||
default_response_class=JavaJSONResponse,
|
||||
lifespan=lifespan,
|
||||
)
|
||||
|
||||
app.add_middleware(AuthenticationMiddleware)
|
||||
app.add_middleware(
|
||||
CORSMiddleware,
|
||||
allow_origins=[],
|
||||
allow_origin_regex=".*",
|
||||
allow_credentials=True,
|
||||
allow_methods=["GET", "POST", "PUT", "DELETE", "OPTIONS"],
|
||||
allow_headers=["*"],
|
||||
max_age=3600,
|
||||
)
|
||||
|
||||
for router in application_routers():
|
||||
app.include_router(router, prefix=settings.context_path)
|
||||
|
||||
|
||||
@app.get(f"{settings.context_path}/health", include_in_schema=False)
|
||||
async def health() -> JavaJSONResponse:
|
||||
return ok({"status": "UP"})
|
||||
|
||||
|
||||
@app.get(f"{settings.context_path}/health/live", include_in_schema=False)
|
||||
async def liveness() -> JavaJSONResponse:
|
||||
return ok({"status": "UP"})
|
||||
|
||||
|
||||
def upload_storage_ready() -> bool:
|
||||
"""Report whether the non-root API process can traverse and write its upload mount."""
|
||||
|
||||
try:
|
||||
return settings.upload_dir.is_dir() and os.access(
|
||||
settings.upload_dir,
|
||||
os.W_OK | os.X_OK,
|
||||
)
|
||||
except OSError:
|
||||
return False
|
||||
|
||||
|
||||
@app.get(f"{settings.context_path}/health/ready", include_in_schema=False)
|
||||
async def readiness() -> JavaJSONResponse:
|
||||
database, redis, uploads = await database_ping(), await redis_ping(), upload_storage_ready()
|
||||
code = 0 if database and redis and uploads else 503
|
||||
msg = "success" if code == 0 else "dependencies unavailable"
|
||||
return JavaJSONResponse(
|
||||
{
|
||||
"code": code,
|
||||
"msg": msg,
|
||||
"data": {"database": database, "redis": redis, "uploads": uploads},
|
||||
},
|
||||
status_code=200 if code == 0 else 503,
|
||||
)
|
||||
|
||||
|
||||
@app.exception_handler(AppError)
|
||||
async def app_error_handler(request: Request, exc: AppError) -> JavaJSONResponse:
|
||||
return error_response(request, exc.code, exc.message, params=exc.params)
|
||||
|
||||
|
||||
@app.exception_handler(RequestValidationError)
|
||||
async def validation_error_handler(request: Request, exc: RequestValidationError) -> JavaJSONResponse:
|
||||
errors = exc.errors()
|
||||
# Spring only maps MethodArgumentNotValidException (a deserialized JSON
|
||||
# object's @Valid field constraints) to code 10034. Root-body conversion,
|
||||
# missing query parameters and multipart binding failures reach its generic
|
||||
# exception handler and therefore keep the HTTP-200/code-500 envelope.
|
||||
root_body_error = any(tuple(error.get("loc", ())) == ("body",) for error in errors)
|
||||
missing_query = any(
|
||||
error.get("type") == "missing" and tuple(error.get("loc", ()))[:1] == ("query",)
|
||||
for error in errors
|
||||
)
|
||||
multipart_binding_error = any(
|
||||
error.get("type") == "missing"
|
||||
and tuple(error.get("loc", ()))[:1] == ("body",)
|
||||
and any(_matches_path_template(request.url.path, path) for path in _MULTIPART_VALIDATION_PATHS)
|
||||
for error in errors
|
||||
)
|
||||
if root_body_error or missing_query or multipart_binding_error:
|
||||
return error_response(request, ErrorCode.INTERNAL_SERVER_ERROR)
|
||||
first = errors[0] if errors else None
|
||||
detail = _java_required_message(request, errors) or (str(first.get("msg")) if first else None)
|
||||
return error_response(request, ErrorCode.PARAM_VALUE_NULL, detail)
|
||||
|
||||
|
||||
@app.exception_handler(IntegrityError)
|
||||
async def integrity_error_handler(request: Request, _: IntegrityError) -> JavaJSONResponse:
|
||||
return error_response(request, ErrorCode.DB_RECORD_EXISTS)
|
||||
|
||||
|
||||
@app.exception_handler(StarletteHTTPException)
|
||||
async def http_error_handler(request: Request, exc: StarletteHTTPException) -> JavaJSONResponse:
|
||||
if exc.status_code == 404:
|
||||
not_found = message_for(ErrorCode.RESOURCE_NOT_FOUND, request.headers.get("Accept-Language"))
|
||||
return error_response(request, 404, not_found)
|
||||
return error_response(request, exc.status_code, str(exc.detail))
|
||||
|
||||
|
||||
@app.exception_handler(Exception)
|
||||
async def unhandled_error_handler(request: Request, exc: Exception) -> JavaJSONResponse:
|
||||
logger.exception("Unhandled manager-api error", exc_info=exc)
|
||||
return error_response(request, ErrorCode.INTERNAL_SERVER_ERROR)
|
||||
@@ -0,0 +1,665 @@
|
||||
from __future__ import annotations
|
||||
|
||||
# All interpolated SQL fragments are selected from closed column/table allowlists.
|
||||
# ruff: noqa: S608
|
||||
from collections.abc import Mapping, Sequence
|
||||
from datetime import datetime
|
||||
from typing import Any
|
||||
|
||||
from sqlalchemy import bindparam, text
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
|
||||
from app.core.database import Repository
|
||||
|
||||
AGENT_COLUMNS = (
|
||||
"id",
|
||||
"user_id",
|
||||
"agent_code",
|
||||
"agent_name",
|
||||
"asr_model_id",
|
||||
"vad_model_id",
|
||||
"llm_model_id",
|
||||
"slm_model_id",
|
||||
"vllm_model_id",
|
||||
"tts_model_id",
|
||||
"tts_voice_id",
|
||||
"tts_language",
|
||||
"tts_volume",
|
||||
"tts_rate",
|
||||
"tts_pitch",
|
||||
"mem_model_id",
|
||||
"intent_model_id",
|
||||
"chat_history_conf",
|
||||
"system_prompt",
|
||||
"summary_memory",
|
||||
"lang_code",
|
||||
"language",
|
||||
"sort",
|
||||
"creator",
|
||||
"created_at",
|
||||
"updater",
|
||||
"updated_at",
|
||||
)
|
||||
AGENT_MUTABLE_COLUMNS = frozenset(AGENT_COLUMNS) - {"id", "user_id", "creator", "created_at"}
|
||||
TEMPLATE_COLUMNS = (
|
||||
"id",
|
||||
"agent_code",
|
||||
"agent_name",
|
||||
"asr_model_id",
|
||||
"vad_model_id",
|
||||
"llm_model_id",
|
||||
"vllm_model_id",
|
||||
"tts_model_id",
|
||||
"tts_voice_id",
|
||||
"tts_language",
|
||||
"tts_volume",
|
||||
"tts_rate",
|
||||
"tts_pitch",
|
||||
"mem_model_id",
|
||||
"intent_model_id",
|
||||
"chat_history_conf",
|
||||
"system_prompt",
|
||||
"summary_memory",
|
||||
"lang_code",
|
||||
"language",
|
||||
"sort",
|
||||
"creator",
|
||||
"created_at",
|
||||
"updater",
|
||||
"updated_at",
|
||||
)
|
||||
|
||||
|
||||
class AgentRepository(Repository):
|
||||
def __init__(self, session: AsyncSession):
|
||||
super().__init__(session)
|
||||
|
||||
@property
|
||||
def is_sqlite(self) -> bool:
|
||||
bind = self.session.get_bind()
|
||||
return bool(bind is not None and bind.dialect.name == "sqlite")
|
||||
|
||||
async def get_agent(self, agent_id: str, *, for_update: bool = False) -> dict[str, Any] | None:
|
||||
suffix = "" if self.is_sqlite or not for_update else " FOR UPDATE"
|
||||
return await self.fetch_one(
|
||||
f"SELECT {', '.join(AGENT_COLUMNS)} FROM ai_agent WHERE id=:id{suffix}", {"id": agent_id}
|
||||
)
|
||||
|
||||
async def check_agent_owner(self, agent_id: str, user_id: int, *, super_admin: bool) -> bool:
|
||||
if super_admin:
|
||||
return bool(await self.scalar("SELECT 1 FROM ai_agent WHERE id=:id LIMIT 1", {"id": agent_id}))
|
||||
return bool(
|
||||
await self.scalar(
|
||||
"SELECT 1 FROM ai_agent WHERE id=:id AND user_id=:user_id LIMIT 1",
|
||||
{"id": agent_id, "user_id": user_id},
|
||||
)
|
||||
)
|
||||
|
||||
async def list_user_agents(self, user_id: int, keyword: str | None) -> list[dict[str, Any]]:
|
||||
params: dict[str, Any] = {"user_id": user_id}
|
||||
where = "a.user_id=:user_id"
|
||||
if keyword is not None and keyword.strip():
|
||||
params["keyword"] = f"%{keyword}%"
|
||||
where += (
|
||||
" AND (a.agent_name LIKE :keyword"
|
||||
" OR EXISTS (SELECT 1 FROM ai_device d0 WHERE d0.agent_id=a.id"
|
||||
" AND d0.user_id=:user_id AND d0.mac_address LIKE :keyword)"
|
||||
" OR EXISTS (SELECT 1 FROM ai_agent_tag_relation tr0"
|
||||
" JOIN ai_agent_tag t0 ON t0.id=tr0.tag_id"
|
||||
" WHERE tr0.agent_id=a.id AND t0.deleted=0 AND t0.tag_name LIKE :keyword))"
|
||||
)
|
||||
return await self.fetch_all(
|
||||
"SELECT a.*, mt.model_name AS tts_model_name, ml.model_name AS llm_model_name,"
|
||||
" mv.model_name AS vllm_model_name, COALESCE(tv.name, vc.name) AS tts_voice_name,"
|
||||
" (SELECT MAX(d.last_connected_at) FROM ai_device d WHERE d.agent_id=a.id) AS last_connected_at,"
|
||||
" (SELECT COUNT(*) FROM ai_device d WHERE d.agent_id=a.id) AS device_count"
|
||||
" FROM ai_agent a"
|
||||
" LEFT JOIN ai_model_config mt ON mt.id=a.tts_model_id"
|
||||
" LEFT JOIN ai_model_config ml ON ml.id=a.llm_model_id"
|
||||
" LEFT JOIN ai_model_config mv ON mv.id=a.vllm_model_id"
|
||||
" LEFT JOIN ai_tts_voice tv ON tv.id=a.tts_voice_id"
|
||||
" LEFT JOIN ai_voice_clone vc ON vc.id=a.tts_voice_id"
|
||||
f" WHERE {where} ORDER BY a.created_at DESC",
|
||||
params,
|
||||
)
|
||||
|
||||
async def list_admin_agents(
|
||||
self, page: int, limit: int, order_field: str, ascending: bool
|
||||
) -> tuple[list[dict[str, Any]], int]:
|
||||
allowed = {"agent_name", "created_at", "updated_at", "sort", "id"}
|
||||
selected = order_field if order_field in allowed else "agent_name"
|
||||
direction = "ASC" if ascending else "DESC"
|
||||
total = int(await self.scalar("SELECT COUNT(*) FROM ai_agent") or 0)
|
||||
query = (
|
||||
f"SELECT {', '.join(AGENT_COLUMNS)} FROM ai_agent "
|
||||
f"ORDER BY {selected} {direction} LIMIT :limit OFFSET :offset"
|
||||
)
|
||||
rows = await self.fetch_all(
|
||||
query,
|
||||
{"limit": limit, "offset": (page - 1) * limit},
|
||||
)
|
||||
return rows, total
|
||||
|
||||
async def insert_agent(self, values: Mapping[str, Any]) -> int:
|
||||
columns = [column for column in AGENT_COLUMNS if column in values]
|
||||
placeholders = ", ".join(f":{column}" for column in columns)
|
||||
return await self.execute(
|
||||
f"INSERT INTO ai_agent ({', '.join(columns)}) VALUES ({placeholders})",
|
||||
{column: values[column] for column in columns},
|
||||
)
|
||||
|
||||
async def update_agent(self, agent_id: str, values: Mapping[str, Any], *, include_null: bool = False) -> int:
|
||||
selected = {
|
||||
key: value
|
||||
for key, value in values.items()
|
||||
if key in AGENT_MUTABLE_COLUMNS and (include_null or value is not None)
|
||||
}
|
||||
if not selected:
|
||||
return 0
|
||||
assignments = ", ".join(f"{column}=:{column}" for column in selected)
|
||||
return await self.execute(
|
||||
f"UPDATE ai_agent SET {assignments} WHERE id=:agent_id",
|
||||
{**selected, "agent_id": agent_id},
|
||||
)
|
||||
|
||||
async def get_agent_plugins(self, agent_id: str) -> list[dict[str, Any]]:
|
||||
return await self.fetch_all(
|
||||
"SELECT m.id,m.agent_id,m.plugin_id,m.param_info,p.provider_code"
|
||||
" FROM ai_agent_plugin_mapping m LEFT JOIN ai_model_provider p ON p.id=m.plugin_id"
|
||||
" WHERE m.agent_id=:agent_id ORDER BY m.id ASC",
|
||||
{"agent_id": agent_id},
|
||||
)
|
||||
|
||||
async def replace_plugins(self, agent_id: str, plugins: Sequence[Mapping[str, Any]]) -> None:
|
||||
existing = await self.fetch_all(
|
||||
"SELECT id,plugin_id FROM ai_agent_plugin_mapping WHERE agent_id=:agent_id",
|
||||
{"agent_id": agent_id},
|
||||
)
|
||||
by_plugin = {str(row["plugin_id"]): int(row["id"]) for row in existing}
|
||||
incoming = {str(item.get("plugin_id") or "") for item in plugins}
|
||||
remove_ids = [int(row["id"]) for row in existing if str(row["plugin_id"]) not in incoming]
|
||||
if remove_ids:
|
||||
statement = text("DELETE FROM ai_agent_plugin_mapping WHERE id IN :ids").bindparams(
|
||||
bindparam("ids", expanding=True)
|
||||
)
|
||||
await self.execute(statement, {"ids": remove_ids})
|
||||
for item in plugins:
|
||||
plugin_id = str(item.get("plugin_id") or "")
|
||||
params = {"agent_id": agent_id, "plugin_id": plugin_id, "param_info": item.get("param_info") or "{}"}
|
||||
if plugin_id in by_plugin:
|
||||
await self.execute(
|
||||
"UPDATE ai_agent_plugin_mapping SET param_info=:param_info WHERE id=:id",
|
||||
{"id": by_plugin[plugin_id], **params},
|
||||
)
|
||||
else:
|
||||
await self.execute(
|
||||
"INSERT INTO ai_agent_plugin_mapping (id,agent_id,plugin_id,param_info)"
|
||||
" VALUES (:id,:agent_id,:plugin_id,:param_info)",
|
||||
{"id": int(item["id"]), **params},
|
||||
)
|
||||
|
||||
async def delete_plugins(self, agent_id: str) -> int:
|
||||
return await self.execute("DELETE FROM ai_agent_plugin_mapping WHERE agent_id=:id", {"id": agent_id})
|
||||
|
||||
async def get_context_provider(self, agent_id: str) -> dict[str, Any] | None:
|
||||
return await self.fetch_one(
|
||||
"SELECT id,agent_id,context_providers,creator,created_at,updater,updated_at"
|
||||
" FROM ai_agent_context_provider WHERE agent_id=:id LIMIT 1",
|
||||
{"id": agent_id},
|
||||
)
|
||||
|
||||
async def upsert_context_provider(self, agent_id: str, encoded: str, new_id: str) -> None:
|
||||
existing = await self.get_context_provider(agent_id)
|
||||
if existing:
|
||||
await self.execute(
|
||||
"UPDATE ai_agent_context_provider SET context_providers=:value WHERE id=:id",
|
||||
{"value": encoded, "id": existing["id"]},
|
||||
)
|
||||
else:
|
||||
await self.execute(
|
||||
"INSERT INTO ai_agent_context_provider (id,agent_id,context_providers)"
|
||||
" VALUES (:id,:agent_id,:value)",
|
||||
{"id": new_id, "agent_id": agent_id, "value": encoded},
|
||||
)
|
||||
|
||||
async def get_correct_word_ids(self, agent_id: str) -> list[str]:
|
||||
rows = await self.fetch_all(
|
||||
"SELECT file_id FROM ai_agent_correct_word_mapping WHERE agent_id=:id", {"id": agent_id}
|
||||
)
|
||||
return [str(row["file_id"]) for row in rows]
|
||||
|
||||
async def replace_correct_words(
|
||||
self, agent_id: str, file_ids: Sequence[str], user_id: int, now: datetime, ids: Sequence[str]
|
||||
) -> None:
|
||||
await self.execute("DELETE FROM ai_agent_correct_word_mapping WHERE agent_id=:id", {"id": agent_id})
|
||||
await self.execute_many(
|
||||
"INSERT INTO ai_agent_correct_word_mapping"
|
||||
" (id,agent_id,file_id,creator,created_at,updater,updated_at)"
|
||||
" VALUES (:id,:agent_id,:file_id,:user_id,:now,:user_id,:now)",
|
||||
[
|
||||
{"id": mapping_id, "agent_id": agent_id, "file_id": file_id, "user_id": user_id, "now": now}
|
||||
for mapping_id, file_id in zip(ids, file_ids, strict=True)
|
||||
],
|
||||
)
|
||||
|
||||
async def get_agent_tags(self, agent_id: str) -> list[dict[str, Any]]:
|
||||
return await self.fetch_all(
|
||||
"SELECT t.id,t.tag_name,t.sort,r.sort AS relation_sort"
|
||||
" FROM ai_agent_tag t JOIN ai_agent_tag_relation r ON t.id=r.tag_id"
|
||||
" WHERE r.agent_id=:id AND t.deleted=0 ORDER BY r.sort ASC,r.created_at ASC",
|
||||
{"id": agent_id},
|
||||
)
|
||||
|
||||
async def get_tags_for_agents(self, agent_ids: Sequence[str]) -> list[dict[str, Any]]:
|
||||
if not agent_ids:
|
||||
return []
|
||||
statement = text(
|
||||
"SELECT t.id,t.tag_name,r.agent_id,r.sort AS relation_sort"
|
||||
" FROM ai_agent_tag t JOIN ai_agent_tag_relation r ON t.id=r.tag_id"
|
||||
" WHERE r.agent_id IN :ids AND t.deleted=0 ORDER BY r.sort ASC,r.created_at ASC"
|
||||
).bindparams(bindparam("ids", expanding=True))
|
||||
return await self.fetch_all(statement, {"ids": list(agent_ids)})
|
||||
|
||||
async def list_tags(self) -> list[dict[str, Any]]:
|
||||
return await self.fetch_all("SELECT id,tag_name,sort FROM ai_agent_tag WHERE deleted=0 ORDER BY sort ASC")
|
||||
|
||||
async def get_tag(self, tag_id: str) -> dict[str, Any] | None:
|
||||
return await self.fetch_one("SELECT * FROM ai_agent_tag WHERE id=:id", {"id": tag_id})
|
||||
|
||||
async def find_active_tag_by_name(self, tag_name: str) -> dict[str, Any] | None:
|
||||
return await self.fetch_one(
|
||||
"SELECT * FROM ai_agent_tag WHERE tag_name=:name AND deleted=0 LIMIT 1", {"name": tag_name}
|
||||
)
|
||||
|
||||
async def find_any_tag_by_name(self, tag_name: str) -> dict[str, Any] | None:
|
||||
return await self.fetch_one("SELECT * FROM ai_agent_tag WHERE tag_name=:name LIMIT 1", {"name": tag_name})
|
||||
|
||||
async def insert_tag(self, values: Mapping[str, Any]) -> int:
|
||||
return await self.execute(
|
||||
"INSERT INTO ai_agent_tag"
|
||||
" (id,tag_name,sort,deleted,creator,created_at,updater,updated_at)"
|
||||
" VALUES (:id,:tag_name,:sort,:deleted,:creator,:created_at,:updater,:updated_at)",
|
||||
values,
|
||||
)
|
||||
|
||||
async def soft_delete_tag(self, tag_id: str, now: datetime) -> int:
|
||||
return await self.execute(
|
||||
"UPDATE ai_agent_tag SET deleted=1,updated_at=:now WHERE id=:id",
|
||||
{"id": tag_id, "now": now},
|
||||
)
|
||||
|
||||
async def replace_tag_relations(self, agent_id: str, relations: Sequence[Mapping[str, Any]]) -> None:
|
||||
await self.execute("DELETE FROM ai_agent_tag_relation WHERE agent_id=:id", {"id": agent_id})
|
||||
await self.execute_many(
|
||||
"INSERT INTO ai_agent_tag_relation"
|
||||
" (id,agent_id,tag_id,sort,creator,created_at,updater,updated_at)"
|
||||
" VALUES (:id,:agent_id,:tag_id,:sort,:creator,:created_at,:updater,:updated_at)",
|
||||
relations,
|
||||
)
|
||||
|
||||
async def get_model_config(self, model_id: str | None) -> dict[str, Any] | None:
|
||||
if not model_id:
|
||||
return None
|
||||
return await self.fetch_one("SELECT * FROM ai_model_config WHERE id=:id", {"id": model_id})
|
||||
|
||||
async def get_default_llm_config(self) -> dict[str, Any] | None:
|
||||
return await self.fetch_one(
|
||||
"SELECT * FROM ai_model_config WHERE model_type='LLM' AND is_enabled=1"
|
||||
" ORDER BY is_default DESC,sort ASC LIMIT 1"
|
||||
)
|
||||
|
||||
async def get_model_provider(self, provider_id: str) -> dict[str, Any] | None:
|
||||
return await self.fetch_one("SELECT * FROM ai_model_provider WHERE id=:id", {"id": provider_id})
|
||||
|
||||
async def get_timbre(self, timbre_id: str | None) -> dict[str, Any] | None:
|
||||
if not timbre_id:
|
||||
return None
|
||||
row = await self.fetch_one("SELECT * FROM ai_tts_voice WHERE id=:id", {"id": timbre_id})
|
||||
if row is None:
|
||||
row = await self.fetch_one("SELECT * FROM ai_voice_clone WHERE id=:id", {"id": timbre_id})
|
||||
return row
|
||||
|
||||
async def find_timbre_by_voice_code(self, model_id: str, voice_code: str) -> dict[str, Any] | None:
|
||||
return await self.fetch_one(
|
||||
"SELECT * FROM ai_tts_voice WHERE tts_model_id=:model_id AND tts_voice=:voice LIMIT 1",
|
||||
{"model_id": model_id, "voice": voice_code},
|
||||
)
|
||||
|
||||
async def get_device_by_mac(self, mac_address: str) -> dict[str, Any] | None:
|
||||
return await self.fetch_one(
|
||||
"SELECT * FROM ai_device WHERE mac_address=:mac ORDER BY id DESC LIMIT 1", {"mac": mac_address}
|
||||
)
|
||||
|
||||
async def get_agent_by_device_mac(self, mac_address: str) -> dict[str, Any] | None:
|
||||
return await self.fetch_one(
|
||||
f"SELECT {', '.join('a.' + column for column in AGENT_COLUMNS)}"
|
||||
" FROM ai_device d LEFT JOIN ai_agent a ON d.agent_id=a.id"
|
||||
" WHERE d.mac_address=:mac ORDER BY d.id DESC LIMIT 1",
|
||||
{"mac": mac_address},
|
||||
)
|
||||
|
||||
async def update_device_connection(self, device_id: str, now: datetime) -> int:
|
||||
return await self.execute(
|
||||
"UPDATE ai_device SET last_connected_at=:now WHERE id=:id", {"id": device_id, "now": now}
|
||||
)
|
||||
|
||||
async def insert_chat_audio(self, audio_id: str, audio: bytes) -> int:
|
||||
return await self.execute(
|
||||
"INSERT INTO ai_agent_chat_audio (id,audio) VALUES (:id,:audio)", {"id": audio_id, "audio": audio}
|
||||
)
|
||||
|
||||
async def get_chat_audio(self, audio_id: str) -> bytes | None:
|
||||
value = await self.scalar("SELECT audio FROM ai_agent_chat_audio WHERE id=:id", {"id": audio_id})
|
||||
return bytes(value) if value is not None else None
|
||||
|
||||
async def insert_chat_history(self, values: Mapping[str, Any]) -> int:
|
||||
return await self.execute(
|
||||
"INSERT INTO ai_agent_chat_history"
|
||||
" (mac_address,agent_id,session_id,chat_type,content,audio_id,created_at)"
|
||||
" VALUES (:mac_address,:agent_id,:session_id,:chat_type,:content,:audio_id,:created_at)",
|
||||
values,
|
||||
)
|
||||
|
||||
async def get_session_agent_id(self, session_id: str) -> str | None:
|
||||
value = await self.scalar(
|
||||
"SELECT agent_id FROM ai_agent_chat_history WHERE session_id=:id LIMIT 1", {"id": session_id}
|
||||
)
|
||||
return str(value) if value is not None else None
|
||||
|
||||
async def get_audio_agent_id(self, audio_id: str) -> str | None:
|
||||
value = await self.scalar(
|
||||
"SELECT agent_id FROM ai_agent_chat_history WHERE audio_id=:id LIMIT 1", {"id": audio_id}
|
||||
)
|
||||
return str(value) if value is not None else None
|
||||
|
||||
async def is_audio_owned(self, audio_id: str, agent_id: str) -> bool:
|
||||
count = await self.scalar(
|
||||
"SELECT COUNT(*) FROM ai_agent_chat_history WHERE audio_id=:audio_id AND agent_id=:agent_id",
|
||||
{"audio_id": audio_id, "agent_id": agent_id},
|
||||
)
|
||||
return int(count or 0) == 1
|
||||
|
||||
async def get_audio_content(self, audio_id: str) -> str | None:
|
||||
value = await self.scalar(
|
||||
"SELECT content FROM ai_agent_chat_history WHERE audio_id=:id LIMIT 1", {"id": audio_id}
|
||||
)
|
||||
return str(value) if value is not None else None
|
||||
|
||||
async def get_chat_history(self, agent_id: str, session_id: str) -> list[dict[str, Any]]:
|
||||
return await self.fetch_all(
|
||||
"SELECT created_at,chat_type,content,audio_id,mac_address"
|
||||
" FROM ai_agent_chat_history WHERE agent_id=:agent_id AND session_id=:session_id"
|
||||
" ORDER BY created_at ASC",
|
||||
{"agent_id": agent_id, "session_id": session_id},
|
||||
)
|
||||
|
||||
async def get_recent_user_history(self, agent_id: str) -> list[dict[str, Any]]:
|
||||
return await self.fetch_all(
|
||||
"SELECT content,audio_id FROM ai_agent_chat_history"
|
||||
" WHERE agent_id=:id AND chat_type=1 AND audio_id IS NOT NULL ORDER BY id DESC LIMIT 50",
|
||||
{"id": agent_id},
|
||||
)
|
||||
|
||||
async def list_sessions(self, agent_id: str, page: int, limit: int) -> tuple[list[dict[str, Any]], int]:
|
||||
total = int(
|
||||
await self.scalar(
|
||||
"SELECT COUNT(*) FROM (SELECT session_id FROM ai_agent_chat_history"
|
||||
" WHERE agent_id=:id GROUP BY session_id) sessions",
|
||||
{"id": agent_id},
|
||||
)
|
||||
or 0
|
||||
)
|
||||
rows = await self.fetch_all(
|
||||
"SELECT h.session_id,MAX(h.created_at) AS created_at,COUNT(*) AS chat_count,"
|
||||
" (SELECT t.title FROM ai_agent_chat_title t WHERE t.session_id=h.session_id LIMIT 1) AS title"
|
||||
" FROM ai_agent_chat_history h WHERE h.agent_id=:id GROUP BY h.session_id"
|
||||
" ORDER BY created_at DESC LIMIT :limit OFFSET :offset",
|
||||
{"id": agent_id, "limit": limit, "offset": (page - 1) * limit},
|
||||
)
|
||||
return rows, total
|
||||
|
||||
async def upsert_chat_title(self, session_id: str, title: str, now: datetime, title_id: str) -> None:
|
||||
existing = await self.fetch_one(
|
||||
"SELECT id FROM ai_agent_chat_title WHERE session_id=:session_id LIMIT 1", {"session_id": session_id}
|
||||
)
|
||||
if existing:
|
||||
await self.execute(
|
||||
"UPDATE ai_agent_chat_title SET title=:title,updated_at=:now WHERE id=:id",
|
||||
{"id": existing["id"], "title": title, "now": now},
|
||||
)
|
||||
else:
|
||||
await self.execute(
|
||||
"INSERT INTO ai_agent_chat_title (id,session_id,title,created_at,updated_at)"
|
||||
" VALUES (:id,:session_id,:title,:now,:now)",
|
||||
{"id": title_id, "session_id": session_id, "title": title, "now": now},
|
||||
)
|
||||
|
||||
async def delete_chat_history(self, agent_id: str, *, delete_audio: bool, delete_text: bool) -> None:
|
||||
if delete_audio:
|
||||
ids = await self.fetch_all(
|
||||
"SELECT DISTINCT audio_id FROM ai_agent_chat_history WHERE agent_id=:id AND audio_id IS NOT NULL",
|
||||
{"id": agent_id},
|
||||
)
|
||||
audio_ids = [str(row["audio_id"]) for row in ids]
|
||||
for offset in range(0, len(audio_ids), 1000):
|
||||
batch = audio_ids[offset : offset + 1000]
|
||||
statement = text("DELETE FROM ai_agent_chat_audio WHERE id IN :ids").bindparams(
|
||||
bindparam("ids", expanding=True)
|
||||
)
|
||||
await self.execute(statement, {"ids": batch})
|
||||
if delete_audio and not delete_text:
|
||||
await self.execute("UPDATE ai_agent_chat_history SET audio_id=NULL WHERE agent_id=:id", {"id": agent_id})
|
||||
if delete_text:
|
||||
await self.execute("DELETE FROM ai_agent_chat_history WHERE agent_id=:id", {"id": agent_id})
|
||||
|
||||
async def delete_agent_cascade(self, agent_id: str) -> None:
|
||||
devices = await self.fetch_all("SELECT mac_address FROM ai_device WHERE agent_id=:id", {"id": agent_id})
|
||||
macs = [str(row["mac_address"]) for row in devices if row.get("mac_address") is not None]
|
||||
await self.execute("DELETE FROM ai_device WHERE agent_id=:id", {"id": agent_id})
|
||||
if macs:
|
||||
statement = text(
|
||||
"DELETE FROM ai_device_address_book WHERE mac_address IN :macs OR target_mac IN :targets"
|
||||
).bindparams(bindparam("macs", expanding=True), bindparam("targets", expanding=True))
|
||||
await self.execute(statement, {"macs": macs, "targets": macs})
|
||||
await self.delete_chat_history(agent_id, delete_audio=True, delete_text=True)
|
||||
for table in (
|
||||
"ai_agent_plugin_mapping",
|
||||
"ai_agent_context_provider",
|
||||
"ai_agent_correct_word_mapping",
|
||||
"ai_agent_tag_relation",
|
||||
"ai_agent_snapshot",
|
||||
):
|
||||
await self.execute(f"DELETE FROM {table} WHERE agent_id=:id", {"id": agent_id})
|
||||
await self.execute("DELETE FROM ai_agent WHERE id=:id", {"id": agent_id})
|
||||
|
||||
async def list_templates(
|
||||
self, *, name: str | None = None, page: int | None = None, limit: int | None = None
|
||||
) -> tuple[list[dict[str, Any]], int]:
|
||||
params: dict[str, Any] = {}
|
||||
where = ""
|
||||
if name:
|
||||
where = " WHERE agent_name LIKE :name"
|
||||
params["name"] = f"%{name}%"
|
||||
total = int(await self.scalar(f"SELECT COUNT(*) FROM ai_agent_template{where}", params) or 0)
|
||||
paging = ""
|
||||
if page is not None and limit is not None:
|
||||
params.update(limit=limit, offset=(page - 1) * limit)
|
||||
paging = " LIMIT :limit OFFSET :offset"
|
||||
rows = await self.fetch_all(f"SELECT * FROM ai_agent_template{where} ORDER BY sort ASC{paging}", params)
|
||||
return rows, total
|
||||
|
||||
async def get_template(self, template_id: str) -> dict[str, Any] | None:
|
||||
return await self.fetch_one("SELECT * FROM ai_agent_template WHERE id=:id", {"id": template_id})
|
||||
|
||||
async def get_default_template(self) -> dict[str, Any] | None:
|
||||
return await self.fetch_one("SELECT * FROM ai_agent_template ORDER BY sort ASC LIMIT 1")
|
||||
|
||||
async def next_template_sort(self) -> int:
|
||||
rows = await self.fetch_all("SELECT sort FROM ai_agent_template WHERE sort IS NOT NULL ORDER BY sort ASC")
|
||||
expected = 1
|
||||
for row in rows:
|
||||
value = int(row["sort"])
|
||||
if value > expected:
|
||||
return expected
|
||||
expected = value + 1
|
||||
return expected
|
||||
|
||||
async def insert_template(self, values: Mapping[str, Any]) -> int:
|
||||
columns = [column for column in TEMPLATE_COLUMNS if column in values]
|
||||
return await self.execute(
|
||||
f"INSERT INTO ai_agent_template ({', '.join(columns)})"
|
||||
f" VALUES ({', '.join(':' + column for column in columns)})",
|
||||
{column: values[column] for column in columns},
|
||||
)
|
||||
|
||||
async def update_template(self, template_id: str, values: Mapping[str, Any]) -> int:
|
||||
selected = {
|
||||
key: value for key, value in values.items() if key in TEMPLATE_COLUMNS and key != "id" and value is not None
|
||||
}
|
||||
if not selected:
|
||||
return 0
|
||||
return await self.execute(
|
||||
f"UPDATE ai_agent_template SET {', '.join(key + '=:' + key for key in selected)} WHERE id=:id",
|
||||
{**selected, "id": template_id},
|
||||
)
|
||||
|
||||
async def delete_template(self, template_id: str) -> int:
|
||||
return await self.execute("DELETE FROM ai_agent_template WHERE id=:id", {"id": template_id})
|
||||
|
||||
async def reorder_templates(self, deleted_sort: int) -> int:
|
||||
return await self.execute("UPDATE ai_agent_template SET sort=sort-1 WHERE sort>:sort", {"sort": deleted_sort})
|
||||
|
||||
async def delete_templates(self, ids: Sequence[str]) -> int:
|
||||
if not ids:
|
||||
return 0
|
||||
statement = text("DELETE FROM ai_agent_template WHERE id IN :ids").bindparams(bindparam("ids", expanding=True))
|
||||
return await self.execute(statement, {"ids": list(ids)})
|
||||
|
||||
async def list_voiceprints(self, agent_id: str, user_id: int) -> list[dict[str, Any]]:
|
||||
return await self.fetch_all(
|
||||
"SELECT id,audio_id,source_name,introduce,create_date"
|
||||
" FROM ai_agent_voice_print WHERE agent_id=:agent_id AND creator=:user_id",
|
||||
{"agent_id": agent_id, "user_id": user_id},
|
||||
)
|
||||
|
||||
async def list_voiceprint_ids(self, agent_id: str) -> list[str]:
|
||||
rows = await self.fetch_all("SELECT id FROM ai_agent_voice_print WHERE agent_id=:id", {"id": agent_id})
|
||||
return [str(row["id"]) for row in rows]
|
||||
|
||||
async def get_voiceprint(self, voiceprint_id: str, user_id: int | None = None) -> dict[str, Any] | None:
|
||||
where = "id=:id"
|
||||
params: dict[str, Any] = {"id": voiceprint_id}
|
||||
if user_id is not None:
|
||||
where += " AND creator=:user_id"
|
||||
params["user_id"] = user_id
|
||||
return await self.fetch_one(f"SELECT * FROM ai_agent_voice_print WHERE {where} LIMIT 1", params)
|
||||
|
||||
async def insert_voiceprint(self, values: Mapping[str, Any]) -> int:
|
||||
return await self.execute(
|
||||
"INSERT INTO ai_agent_voice_print"
|
||||
" (id,agent_id,audio_id,source_name,introduce,creator,create_date,updater,update_date)"
|
||||
" VALUES (:id,:agent_id,:audio_id,:source_name,:introduce,:creator,:create_date,:updater,:update_date)",
|
||||
values,
|
||||
)
|
||||
|
||||
async def update_voiceprint(self, voiceprint_id: str, user_id: int, values: Mapping[str, Any]) -> int:
|
||||
allowed = {"audio_id", "source_name", "introduce", "updater", "update_date"}
|
||||
selected = {key: value for key, value in values.items() if key in allowed and value is not None}
|
||||
if not selected:
|
||||
return 0
|
||||
return await self.execute(
|
||||
f"UPDATE ai_agent_voice_print SET {', '.join(key + '=:' + key for key in selected)}"
|
||||
" WHERE id=:id AND creator=:user_id",
|
||||
{**selected, "id": voiceprint_id, "user_id": user_id},
|
||||
)
|
||||
|
||||
async def delete_voiceprint(self, voiceprint_id: str, user_id: int) -> int:
|
||||
return await self.execute(
|
||||
"DELETE FROM ai_agent_voice_print WHERE id=:id AND creator=:user_id",
|
||||
{"id": voiceprint_id, "user_id": user_id},
|
||||
)
|
||||
|
||||
async def snapshot_max_version(self, agent_id: str) -> int:
|
||||
return int(
|
||||
await self.scalar(
|
||||
"SELECT COALESCE(MAX(version_no),0) FROM ai_agent_snapshot WHERE agent_id=:id", {"id": agent_id}
|
||||
)
|
||||
or 0
|
||||
)
|
||||
|
||||
async def latest_snapshot(self, agent_id: str) -> dict[str, Any] | None:
|
||||
return await self.fetch_one(
|
||||
"SELECT * FROM ai_agent_snapshot WHERE agent_id=:id ORDER BY version_no DESC LIMIT 1", {"id": agent_id}
|
||||
)
|
||||
|
||||
async def next_snapshot(self, agent_id: str, version_no: int) -> dict[str, Any] | None:
|
||||
return await self.fetch_one(
|
||||
"SELECT * FROM ai_agent_snapshot WHERE agent_id=:id AND version_no>:version"
|
||||
" ORDER BY version_no ASC LIMIT 1",
|
||||
{"id": agent_id, "version": version_no},
|
||||
)
|
||||
|
||||
async def get_snapshot(self, snapshot_id: str) -> dict[str, Any] | None:
|
||||
return await self.fetch_one("SELECT * FROM ai_agent_snapshot WHERE id=:id", {"id": snapshot_id})
|
||||
|
||||
async def list_snapshots(
|
||||
self, agent_id: str, page: int, limit: int, max_version_no: int | None
|
||||
) -> tuple[list[dict[str, Any]], int]:
|
||||
params: dict[str, Any] = {"id": agent_id, "limit": limit, "offset": (page - 1) * limit}
|
||||
extra = ""
|
||||
if max_version_no is not None:
|
||||
extra = " AND version_no<=:max_version"
|
||||
params["max_version"] = max_version_no
|
||||
total = int(await self.scalar(f"SELECT COUNT(*) FROM ai_agent_snapshot WHERE agent_id=:id{extra}", params) or 0)
|
||||
rows = await self.fetch_all(
|
||||
f"SELECT * FROM ai_agent_snapshot WHERE agent_id=:id{extra}"
|
||||
" ORDER BY version_no DESC LIMIT :limit OFFSET :offset",
|
||||
params,
|
||||
)
|
||||
return rows, total
|
||||
|
||||
async def insert_snapshot_next_version(self, values: Mapping[str, Any]) -> int:
|
||||
return await self.execute(
|
||||
"INSERT INTO ai_agent_snapshot"
|
||||
" (id,agent_id,user_id,version_no,snapshot_data,changed_fields,source,"
|
||||
" restore_from_snapshot_id,restore_from_version_no,creator,created_at,redaction_version)"
|
||||
" SELECT :id,:agent_id,:user_id,COALESCE(MAX(version_no),0)+1,:snapshot_data,:changed_fields,:source,"
|
||||
" :restore_from_snapshot_id,:restore_from_version_no,:creator,:created_at,:redaction_version"
|
||||
" FROM ai_agent_snapshot WHERE agent_id=:agent_id",
|
||||
values,
|
||||
)
|
||||
|
||||
async def prune_snapshots(self, agent_id: str, keep: int) -> int:
|
||||
rows = await self.fetch_all(
|
||||
"SELECT id FROM ai_agent_snapshot WHERE agent_id=:id ORDER BY version_no DESC LIMIT :keep",
|
||||
{"id": agent_id, "keep": keep},
|
||||
)
|
||||
retained = [str(row["id"]) for row in rows]
|
||||
if not retained:
|
||||
return 0
|
||||
statement = text("DELETE FROM ai_agent_snapshot WHERE agent_id=:agent_id AND id NOT IN :retained").bindparams(
|
||||
bindparam("retained", expanding=True)
|
||||
)
|
||||
return await self.execute(statement, {"agent_id": agent_id, "retained": retained})
|
||||
|
||||
async def delete_snapshot(self, snapshot_id: str) -> int:
|
||||
return await self.execute("DELETE FROM ai_agent_snapshot WHERE id=:id", {"id": snapshot_id})
|
||||
|
||||
async def list_legacy_snapshots(self, after_id: str | None, limit: int, version: int) -> list[dict[str, Any]]:
|
||||
params: dict[str, Any] = {"version": version, "limit": limit}
|
||||
extra = ""
|
||||
if after_id is not None:
|
||||
extra = " AND id>:after_id"
|
||||
params["after_id"] = after_id
|
||||
return await self.fetch_all(
|
||||
f"SELECT id,snapshot_data,redaction_version FROM ai_agent_snapshot"
|
||||
f" WHERE redaction_version<:version{extra} ORDER BY id ASC LIMIT :limit",
|
||||
params,
|
||||
)
|
||||
|
||||
async def update_redacted_snapshot(self, snapshot_id: str, snapshot_data: str, version: int) -> int:
|
||||
return await self.execute(
|
||||
"UPDATE ai_agent_snapshot SET snapshot_data=:data,redaction_version=:version"
|
||||
" WHERE id=:id AND redaction_version<:version",
|
||||
{"id": snapshot_id, "data": snapshot_data, "version": version},
|
||||
)
|
||||
@@ -0,0 +1,105 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Any
|
||||
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
|
||||
from app.core.database import Repository
|
||||
|
||||
|
||||
class ConfigRepository(Repository):
|
||||
def __init__(self, session: AsyncSession):
|
||||
super().__init__(session)
|
||||
|
||||
async def list_params(self) -> list[dict[str, Any]]:
|
||||
return await self.fetch_all(
|
||||
"SELECT param_code, param_value, value_type FROM sys_params WHERE param_type = 1"
|
||||
)
|
||||
|
||||
async def get_param_value(self, code: str) -> str | None:
|
||||
value = await self.scalar(
|
||||
"SELECT param_value FROM sys_params WHERE param_code = :code LIMIT 1",
|
||||
{"code": code},
|
||||
)
|
||||
return None if value is None else str(value)
|
||||
|
||||
async def get_default_template(self) -> dict[str, Any] | None:
|
||||
return await self.fetch_one(
|
||||
"SELECT id, agent_code, agent_name, asr_model_id, vad_model_id, llm_model_id, "
|
||||
"vllm_model_id, tts_model_id, tts_voice_id, tts_language, tts_volume, tts_rate, tts_pitch, "
|
||||
"mem_model_id, intent_model_id, chat_history_conf, system_prompt, summary_memory, lang_code, "
|
||||
"language, sort FROM ai_agent_template ORDER BY sort ASC LIMIT 1"
|
||||
)
|
||||
|
||||
async def get_device_by_mac(self, mac_address: str) -> dict[str, Any] | None:
|
||||
return await self.fetch_one(
|
||||
"SELECT id, user_id, mac_address, board, agent_id, app_version, auto_update "
|
||||
"FROM ai_device WHERE mac_address = :mac_address LIMIT 1",
|
||||
{"mac_address": mac_address},
|
||||
)
|
||||
|
||||
async def get_agent(self, agent_id: str) -> dict[str, Any] | None:
|
||||
return await self.fetch_one(
|
||||
"SELECT id, user_id, agent_code, agent_name, asr_model_id, vad_model_id, llm_model_id, slm_model_id, "
|
||||
"vllm_model_id, tts_model_id, tts_voice_id, tts_language, tts_volume, tts_rate, tts_pitch, "
|
||||
"mem_model_id, intent_model_id, chat_history_conf, system_prompt, summary_memory, lang_code, language "
|
||||
"FROM ai_agent WHERE id = :id LIMIT 1",
|
||||
{"id": agent_id},
|
||||
)
|
||||
|
||||
async def get_model(self, model_id: str) -> dict[str, Any] | None:
|
||||
return await self.fetch_one(
|
||||
"SELECT id, model_type, model_code, model_name, is_default, is_enabled, config_json, doc_link, "
|
||||
"remark, sort, creator, create_date, updater, update_date "
|
||||
"FROM ai_model_config WHERE id = :id LIMIT 1",
|
||||
{"id": model_id},
|
||||
)
|
||||
|
||||
async def get_timbre(self, timbre_id: str) -> dict[str, Any] | None:
|
||||
return await self.fetch_one(
|
||||
"SELECT id, languages, name, remark, reference_audio, reference_text, sort, tts_model_id, "
|
||||
"tts_voice, voice_demo FROM ai_tts_voice WHERE id = :id LIMIT 1",
|
||||
{"id": timbre_id},
|
||||
)
|
||||
|
||||
async def get_voice_clone(self, clone_id: str) -> dict[str, Any] | None:
|
||||
return await self.fetch_one(
|
||||
"SELECT id, name, model_id, voice_id, languages, user_id, train_status, train_error "
|
||||
"FROM ai_voice_clone WHERE id = :id LIMIT 1",
|
||||
{"id": clone_id},
|
||||
)
|
||||
|
||||
async def get_plugin_mappings(self, agent_id: str) -> list[dict[str, Any]]:
|
||||
return await self.fetch_all(
|
||||
"SELECT m.id, m.agent_id, m.plugin_id, m.param_info, "
|
||||
"(SELECT p.provider_code FROM ai_model_provider p WHERE p.id = m.plugin_id LIMIT 1) AS provider_code "
|
||||
"FROM ai_agent_plugin_mapping m WHERE m.agent_id = :agent_id",
|
||||
{"agent_id": agent_id},
|
||||
)
|
||||
|
||||
async def get_dataset(self, dataset_id: str) -> dict[str, Any] | None:
|
||||
return await self.fetch_one(
|
||||
"SELECT id, dataset_id, rag_model_id, name, description, status "
|
||||
"FROM ai_rag_dataset WHERE id = :id LIMIT 1",
|
||||
{"id": dataset_id},
|
||||
)
|
||||
|
||||
async def get_context_providers(self, agent_id: str) -> Any:
|
||||
return await self.scalar(
|
||||
"SELECT context_providers FROM ai_agent_context_provider WHERE agent_id = :agent_id LIMIT 1",
|
||||
{"agent_id": agent_id},
|
||||
)
|
||||
|
||||
async def get_voiceprints(self, agent_id: str) -> list[dict[str, Any]]:
|
||||
return await self.fetch_all(
|
||||
"SELECT id, agent_id, source_name, introduce, create_date "
|
||||
"FROM ai_agent_voice_print WHERE agent_id = :agent_id ORDER BY create_date ASC",
|
||||
{"agent_id": agent_id},
|
||||
)
|
||||
|
||||
async def get_correct_word_items(self, agent_id: str) -> list[dict[str, Any]]:
|
||||
return await self.fetch_all(
|
||||
"SELECT i.source_word, i.target_word FROM ai_agent_correct_word_mapping m "
|
||||
"JOIN ai_agent_correct_word_item i ON i.file_id = m.file_id WHERE m.agent_id = :agent_id",
|
||||
{"agent_id": agent_id},
|
||||
)
|
||||
@@ -0,0 +1,115 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import uuid
|
||||
from collections.abc import Sequence
|
||||
from typing import Any
|
||||
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
|
||||
from app.core.database import Repository
|
||||
|
||||
|
||||
class CorrectWordRepository(Repository):
|
||||
def __init__(self, session: AsyncSession):
|
||||
super().__init__(session)
|
||||
|
||||
async def name_exists(self, user_id: int, file_name: str, exclude_id: str | None = None) -> bool:
|
||||
return bool(
|
||||
await self.scalar(
|
||||
"SELECT COUNT(*) FROM ai_agent_correct_word_file WHERE creator=:creator AND file_name=:file_name "
|
||||
"AND (:exclude_id IS NULL OR id<>:exclude_id)",
|
||||
{"creator": user_id, "file_name": file_name, "exclude_id": exclude_id},
|
||||
)
|
||||
)
|
||||
|
||||
async def insert_file(self, values: dict[str, Any]) -> None:
|
||||
await self.execute(
|
||||
"INSERT INTO ai_agent_correct_word_file "
|
||||
"(id, file_name, word_count, content, creator, created_at) "
|
||||
"VALUES (:id, :file_name, :word_count, :content, :creator, :now)",
|
||||
values,
|
||||
)
|
||||
|
||||
async def insert_items(self, values: Sequence[dict[str, Any]]) -> None:
|
||||
await self.execute_many(
|
||||
"INSERT INTO ai_agent_correct_word_item (id, file_id, source_word, target_word) "
|
||||
"VALUES (:id, :file_id, :source_word, :target_word)",
|
||||
values,
|
||||
)
|
||||
|
||||
async def get_file(self, file_id: str, *, for_update: bool = False) -> dict[str, Any] | None:
|
||||
suffix = " FOR UPDATE" if for_update and self.session.get_bind().dialect.name != "sqlite" else ""
|
||||
return await self.fetch_one(
|
||||
f"SELECT * FROM ai_agent_correct_word_file WHERE id=:id{suffix}", # noqa: S608
|
||||
{"id": file_id},
|
||||
)
|
||||
|
||||
async def update_file(self, values: dict[str, Any]) -> int:
|
||||
return await self.execute(
|
||||
"UPDATE ai_agent_correct_word_file SET file_name=:file_name, word_count=:word_count, "
|
||||
"content=:content, updater=:updater, updated_at=:now WHERE id=:id",
|
||||
values,
|
||||
)
|
||||
|
||||
async def list_files(
|
||||
self, user_id: int, *, offset: int | None = None, limit: int | None = None
|
||||
) -> tuple[list[dict[str, Any]], int]:
|
||||
total = int(
|
||||
await self.scalar(
|
||||
"SELECT COUNT(*) FROM ai_agent_correct_word_file WHERE creator=:creator", {"creator": user_id}
|
||||
)
|
||||
or 0
|
||||
)
|
||||
if offset is None or limit is None:
|
||||
rows = await self.fetch_all(
|
||||
"SELECT * FROM ai_agent_correct_word_file WHERE creator=:creator ORDER BY created_at DESC",
|
||||
{"creator": user_id},
|
||||
)
|
||||
else:
|
||||
rows = await self.fetch_all(
|
||||
"SELECT * FROM ai_agent_correct_word_file WHERE creator=:creator ORDER BY created_at DESC "
|
||||
"LIMIT :offset, :limit",
|
||||
{"creator": user_id, "offset": offset, "limit": limit},
|
||||
)
|
||||
return rows, total
|
||||
|
||||
async def delete_file_graph(self, file_id: str) -> None:
|
||||
await self.execute("DELETE FROM ai_agent_correct_word_mapping WHERE file_id=:id", {"id": file_id})
|
||||
await self.execute("DELETE FROM ai_agent_correct_word_item WHERE file_id=:id", {"id": file_id})
|
||||
await self.execute("DELETE FROM ai_agent_correct_word_file WHERE id=:id", {"id": file_id})
|
||||
|
||||
async def delete_items(self, file_id: str) -> None:
|
||||
await self.execute("DELETE FROM ai_agent_correct_word_item WHERE file_id=:id", {"id": file_id})
|
||||
|
||||
async def items_for_agent(self, agent_id: str) -> list[dict[str, Any]]:
|
||||
return await self.fetch_all(
|
||||
"SELECT i.source_word, i.target_word FROM ai_agent_correct_word_item i "
|
||||
"JOIN ai_agent_correct_word_mapping m ON m.file_id=i.file_id WHERE m.agent_id=:agent_id",
|
||||
{"agent_id": agent_id},
|
||||
)
|
||||
|
||||
async def file_ids_for_agent(self, agent_id: str) -> list[str]:
|
||||
rows = await self.fetch_all(
|
||||
"SELECT file_id FROM ai_agent_correct_word_mapping WHERE agent_id=:agent_id", {"agent_id": agent_id}
|
||||
)
|
||||
return [str(row["file_id"]) for row in rows]
|
||||
|
||||
async def replace_agent_mappings(
|
||||
self, agent_id: str, file_ids: Sequence[str], user_id: int, now: Any
|
||||
) -> None:
|
||||
await self.execute("DELETE FROM ai_agent_correct_word_mapping WHERE agent_id=:agent_id", {"agent_id": agent_id})
|
||||
await self.execute_many(
|
||||
"INSERT INTO ai_agent_correct_word_mapping "
|
||||
"(id, agent_id, file_id, creator, created_at, updater, updated_at) "
|
||||
"VALUES (:id, :agent_id, :file_id, :user_id, :now, :user_id, :now)",
|
||||
[
|
||||
{
|
||||
"id": uuid.uuid4().hex,
|
||||
"agent_id": agent_id,
|
||||
"file_id": file_id,
|
||||
"user_id": user_id,
|
||||
"now": now,
|
||||
}
|
||||
for file_id in file_ids
|
||||
],
|
||||
)
|
||||
@@ -0,0 +1,294 @@
|
||||
from __future__ import annotations
|
||||
|
||||
# Every interpolated SQL fragment below is a module constant or a service-side allowlist.
|
||||
# ruff: noqa: S608
|
||||
from collections.abc import Mapping, Sequence
|
||||
from datetime import datetime
|
||||
from typing import Any
|
||||
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
|
||||
from app.core.database import Repository
|
||||
|
||||
DEVICE_COLUMNS = (
|
||||
"id, user_id, mac_address, last_connected_at, auto_update, board, alias, "
|
||||
"agent_id, app_version, sort, updater, update_date, creator, create_date"
|
||||
)
|
||||
OTA_COLUMNS = (
|
||||
"id, firmware_name, type, version, size, remark, firmware_path, sort, "
|
||||
"updater, update_date, creator, create_date"
|
||||
)
|
||||
ADDRESS_BOOK_COLUMNS = (
|
||||
"mac_address, target_mac, alias, has_permission, creator, create_date, updater, update_date"
|
||||
)
|
||||
|
||||
|
||||
class DeviceRepository(Repository):
|
||||
def __init__(self, session: AsyncSession):
|
||||
super().__init__(session)
|
||||
|
||||
async def get_device(self, device_id: str) -> dict[str, Any] | None:
|
||||
return await self.fetch_one(
|
||||
f"SELECT {DEVICE_COLUMNS} FROM ai_device WHERE id = :id LIMIT 1",
|
||||
{"id": device_id},
|
||||
)
|
||||
|
||||
async def get_device_by_mac(self, mac_address: str) -> dict[str, Any] | None:
|
||||
return await self.fetch_one(
|
||||
f"SELECT {DEVICE_COLUMNS} FROM ai_device WHERE mac_address = :mac_address LIMIT 1",
|
||||
{"mac_address": mac_address},
|
||||
)
|
||||
|
||||
async def get_user_devices(self, user_id: int, agent_id: str) -> list[dict[str, Any]]:
|
||||
return await self.fetch_all(
|
||||
f"SELECT {DEVICE_COLUMNS} FROM ai_device WHERE user_id = :user_id AND agent_id = :agent_id",
|
||||
{"user_id": user_id, "agent_id": agent_id},
|
||||
)
|
||||
|
||||
async def insert_device(self, values: Mapping[str, Any]) -> None:
|
||||
await self.execute(
|
||||
"INSERT INTO ai_device "
|
||||
"(id, user_id, mac_address, last_connected_at, auto_update, board, alias, agent_id, app_version, "
|
||||
"sort, updater, update_date, creator, create_date) "
|
||||
"VALUES (:id, :user_id, :mac_address, :last_connected_at, :auto_update, :board, :alias, :agent_id, "
|
||||
":app_version, :sort, :updater, :update_date, :creator, :create_date)",
|
||||
values,
|
||||
)
|
||||
|
||||
async def update_device_info(
|
||||
self,
|
||||
device_id: str,
|
||||
*,
|
||||
auto_update: int | None,
|
||||
alias: str | None,
|
||||
updater: int,
|
||||
now: datetime,
|
||||
) -> int:
|
||||
assignments = ["updater = :updater", "update_date = :now"]
|
||||
params: dict[str, Any] = {"id": device_id, "updater": updater, "now": now}
|
||||
if auto_update is not None:
|
||||
assignments.append("auto_update = :auto_update")
|
||||
params["auto_update"] = auto_update
|
||||
if alias is not None:
|
||||
assignments.append("alias = :alias")
|
||||
params["alias"] = alias
|
||||
return await self.execute(
|
||||
f"UPDATE ai_device SET {', '.join(assignments)} WHERE id = :id",
|
||||
params,
|
||||
)
|
||||
|
||||
async def update_connection(
|
||||
self,
|
||||
device_id: str,
|
||||
*,
|
||||
app_version: str | None,
|
||||
now: datetime,
|
||||
) -> int:
|
||||
if app_version is None or not app_version.strip():
|
||||
return await self.execute(
|
||||
"UPDATE ai_device SET last_connected_at = :now WHERE id = :id",
|
||||
{"id": device_id, "now": now},
|
||||
)
|
||||
return await self.execute(
|
||||
"UPDATE ai_device SET last_connected_at = :now, app_version = :app_version WHERE id = :id",
|
||||
{"id": device_id, "now": now, "app_version": app_version},
|
||||
)
|
||||
|
||||
async def delete_device_for_user(self, device_id: str, user_id: int) -> int:
|
||||
return await self.execute(
|
||||
"DELETE FROM ai_device WHERE id = :id AND user_id = :user_id",
|
||||
{"id": device_id, "user_id": user_id},
|
||||
)
|
||||
|
||||
async def get_address_book(self, mac_address: str) -> list[dict[str, Any]]:
|
||||
return await self.fetch_all(
|
||||
f"SELECT {ADDRESS_BOOK_COLUMNS} FROM ai_device_address_book "
|
||||
"WHERE mac_address = :mac_address ORDER BY update_date DESC",
|
||||
{"mac_address": mac_address},
|
||||
)
|
||||
|
||||
async def get_all_address_book(self) -> list[dict[str, Any]]:
|
||||
return await self.fetch_all(f"SELECT {ADDRESS_BOOK_COLUMNS} FROM ai_device_address_book")
|
||||
|
||||
async def get_address_book_record(self, mac_address: str, target_mac: str) -> dict[str, Any] | None:
|
||||
return await self.fetch_one(
|
||||
f"SELECT {ADDRESS_BOOK_COLUMNS} FROM ai_device_address_book "
|
||||
"WHERE mac_address = :mac_address AND target_mac = :target_mac LIMIT 1",
|
||||
{"mac_address": mac_address, "target_mac": target_mac},
|
||||
)
|
||||
|
||||
async def get_aliases(self, mac_address: str) -> list[str]:
|
||||
rows = await self.fetch_all(
|
||||
"SELECT alias FROM ai_device_address_book WHERE mac_address = :mac_address",
|
||||
{"mac_address": mac_address},
|
||||
)
|
||||
return [str(row["alias"]) for row in rows if row.get("alias") not in (None, "")]
|
||||
|
||||
async def insert_address_book(
|
||||
self,
|
||||
*,
|
||||
mac_address: str,
|
||||
target_mac: str,
|
||||
alias: str | None,
|
||||
has_permission: bool | None,
|
||||
actor: int,
|
||||
now: datetime,
|
||||
) -> None:
|
||||
await self.execute(
|
||||
"INSERT INTO ai_device_address_book "
|
||||
"(mac_address, target_mac, alias, has_permission, creator, create_date, updater, update_date) "
|
||||
"VALUES (:mac_address, :target_mac, :alias, :has_permission, :actor, :now, :actor, :now)",
|
||||
{
|
||||
"mac_address": mac_address,
|
||||
"target_mac": target_mac,
|
||||
"alias": alias,
|
||||
"has_permission": has_permission,
|
||||
"actor": actor,
|
||||
"now": now,
|
||||
},
|
||||
)
|
||||
|
||||
async def update_address_alias(
|
||||
self,
|
||||
mac_address: str,
|
||||
target_mac: str,
|
||||
alias: str | None,
|
||||
*,
|
||||
now: datetime,
|
||||
) -> int:
|
||||
return await self.execute(
|
||||
"UPDATE ai_device_address_book SET alias = :alias, update_date = :now "
|
||||
"WHERE mac_address = :mac_address AND target_mac = :target_mac",
|
||||
{
|
||||
"mac_address": mac_address,
|
||||
"target_mac": target_mac,
|
||||
"alias": alias,
|
||||
"now": now,
|
||||
},
|
||||
)
|
||||
|
||||
async def update_address_permission(
|
||||
self,
|
||||
mac_address: str,
|
||||
target_mac: str,
|
||||
has_permission: bool,
|
||||
*,
|
||||
now: datetime,
|
||||
) -> int:
|
||||
return await self.execute(
|
||||
"UPDATE ai_device_address_book "
|
||||
"SET has_permission = :has_permission, update_date = :now "
|
||||
"WHERE mac_address = :mac_address AND target_mac = :target_mac",
|
||||
{
|
||||
"mac_address": mac_address,
|
||||
"target_mac": target_mac,
|
||||
"has_permission": has_permission,
|
||||
"now": now,
|
||||
},
|
||||
)
|
||||
|
||||
async def delete_address_books_for_macs(self, mac_addresses: Sequence[str]) -> int:
|
||||
if not mac_addresses:
|
||||
return 0
|
||||
placeholders = ", ".join(f":mac_{index}" for index in range(len(mac_addresses)))
|
||||
params = {f"mac_{index}": mac for index, mac in enumerate(mac_addresses)}
|
||||
return await self.execute(
|
||||
f"DELETE FROM ai_device_address_book WHERE mac_address IN ({placeholders}) "
|
||||
f"OR target_mac IN ({placeholders})",
|
||||
params,
|
||||
)
|
||||
|
||||
async def count_ota(self, firmware_name: str | None = None) -> int:
|
||||
where = ""
|
||||
params: dict[str, Any] = {}
|
||||
if firmware_name is not None and firmware_name.strip():
|
||||
where = " WHERE firmware_name LIKE :firmware_name"
|
||||
params["firmware_name"] = f"%{firmware_name}%"
|
||||
return int(await self.scalar(f"SELECT COUNT(*) FROM ai_ota{where}", params) or 0)
|
||||
|
||||
async def list_ota(
|
||||
self,
|
||||
*,
|
||||
page: int,
|
||||
limit: int,
|
||||
firmware_name: str | None,
|
||||
order_fields: Sequence[str],
|
||||
ascending: bool,
|
||||
) -> list[dict[str, Any]]:
|
||||
where = ""
|
||||
params: dict[str, Any] = {"limit": limit, "offset": max(page - 1, 0) * limit}
|
||||
if firmware_name is not None and firmware_name.strip():
|
||||
where = " WHERE firmware_name LIKE :firmware_name"
|
||||
params["firmware_name"] = f"%{firmware_name}%"
|
||||
direction = "ASC" if ascending else "DESC"
|
||||
order_by = ", ".join(f"{field} {direction}" for field in order_fields)
|
||||
return await self.fetch_all(
|
||||
f"SELECT {OTA_COLUMNS} FROM ai_ota{where} ORDER BY {order_by} LIMIT :limit OFFSET :offset",
|
||||
params,
|
||||
)
|
||||
|
||||
async def get_ota(self, ota_id: str) -> dict[str, Any] | None:
|
||||
return await self.fetch_one(
|
||||
f"SELECT {OTA_COLUMNS} FROM ai_ota WHERE id = :id LIMIT 1",
|
||||
{"id": ota_id},
|
||||
)
|
||||
|
||||
async def get_first_ota_by_type(self, ota_type: str) -> dict[str, Any] | None:
|
||||
return await self.fetch_one(
|
||||
f"SELECT {OTA_COLUMNS} FROM ai_ota WHERE type = :type LIMIT 1",
|
||||
{"type": ota_type},
|
||||
)
|
||||
|
||||
async def get_latest_ota(self, ota_type: str) -> dict[str, Any] | None:
|
||||
return await self.fetch_one(
|
||||
f"SELECT {OTA_COLUMNS} FROM ai_ota WHERE type = :type ORDER BY update_date DESC LIMIT 1",
|
||||
{"type": ota_type},
|
||||
)
|
||||
|
||||
async def count_duplicate_ota(self, *, ota_id: str, ota_type: str | None, version: str | None) -> int:
|
||||
return int(
|
||||
await self.scalar(
|
||||
"SELECT COUNT(*) FROM ai_ota WHERE type = :type AND version = :version AND id <> :id",
|
||||
{"id": ota_id, "type": ota_type, "version": version},
|
||||
)
|
||||
or 0
|
||||
)
|
||||
|
||||
async def insert_ota(self, values: Mapping[str, Any]) -> None:
|
||||
await self.execute(
|
||||
"INSERT INTO ai_ota "
|
||||
"(id, firmware_name, type, version, size, remark, firmware_path, sort, updater, update_date, creator, "
|
||||
"create_date) VALUES (:id, :firmware_name, :type, :version, :size, :remark, :firmware_path, :sort, "
|
||||
":updater, :update_date, :creator, :create_date)",
|
||||
values,
|
||||
)
|
||||
|
||||
async def update_ota(self, ota_id: str, values: Mapping[str, Any]) -> int:
|
||||
allowed = {
|
||||
"firmware_name",
|
||||
"type",
|
||||
"version",
|
||||
"size",
|
||||
"remark",
|
||||
"firmware_path",
|
||||
"sort",
|
||||
"updater",
|
||||
"update_date",
|
||||
"creator",
|
||||
"create_date",
|
||||
}
|
||||
selected = {key: value for key, value in values.items() if key in allowed and value is not None}
|
||||
if not selected:
|
||||
return 0
|
||||
assignments = ", ".join(f"{key} = :{key}" for key in selected)
|
||||
return await self.execute(
|
||||
f"UPDATE ai_ota SET {assignments} WHERE id = :id",
|
||||
{"id": ota_id, **selected},
|
||||
)
|
||||
|
||||
async def delete_ota(self, ids: Sequence[str]) -> int:
|
||||
if not ids:
|
||||
return 0
|
||||
placeholders = ", ".join(f":id_{index}" for index in range(len(ids)))
|
||||
params = {f"id_{index}": value for index, value in enumerate(ids)}
|
||||
return await self.execute(f"DELETE FROM ai_ota WHERE id IN ({placeholders})", params)
|
||||
@@ -0,0 +1,344 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
import uuid
|
||||
from collections.abc import Sequence
|
||||
from datetime import datetime
|
||||
from typing import Any
|
||||
from zoneinfo import ZoneInfo
|
||||
|
||||
from sqlalchemy import bindparam, text
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
|
||||
from app.core.config import get_settings
|
||||
from app.core.database import Repository
|
||||
|
||||
|
||||
class KnowledgeRepository(Repository):
|
||||
def __init__(self, session: AsyncSession):
|
||||
super().__init__(session)
|
||||
|
||||
async def dataset_page(
|
||||
self, user_id: int, name: str | None, offset: int, limit: int
|
||||
) -> tuple[list[dict[str, Any]], int]:
|
||||
where = (
|
||||
"WHERE creator=:creator AND (:name IS NULL OR :name='' OR name LIKE CONCAT('%', :name, '%'))"
|
||||
)
|
||||
params = {"creator": user_id, "name": name, "offset": offset, "limit": limit}
|
||||
total = int(await self.scalar(f"SELECT COUNT(*) FROM ai_rag_dataset {where}", params) or 0) # noqa: S608
|
||||
rows = await self.fetch_all(
|
||||
f"SELECT * FROM ai_rag_dataset {where} ORDER BY created_at DESC LIMIT :offset, :limit", # noqa: S608
|
||||
params,
|
||||
)
|
||||
return rows, total
|
||||
|
||||
async def get_dataset(self, identifier: str, *, for_update: bool = False) -> dict[str, Any] | None:
|
||||
suffix = " FOR UPDATE" if for_update and self.session.get_bind().dialect.name != "sqlite" else ""
|
||||
return await self.fetch_one(
|
||||
f"SELECT * FROM ai_rag_dataset WHERE dataset_id=:id OR id=:id LIMIT 1{suffix}", # noqa: S608
|
||||
{"id": identifier},
|
||||
)
|
||||
|
||||
async def datasets_by_ids(self, identifiers: Sequence[str]) -> list[dict[str, Any]]:
|
||||
if not identifiers:
|
||||
return []
|
||||
statement = text("SELECT * FROM ai_rag_dataset WHERE dataset_id IN :ids OR id IN :ids").bindparams(
|
||||
bindparam("ids", expanding=True)
|
||||
)
|
||||
result = await self.session.execute(statement, {"ids": list(identifiers)})
|
||||
return [dict(row) for row in result.mappings().all()]
|
||||
|
||||
async def duplicate_dataset_name(self, user_id: int, name: str, exclude_id: str | None = None) -> bool:
|
||||
return bool(
|
||||
await self.scalar(
|
||||
"SELECT COUNT(*) FROM ai_rag_dataset WHERE creator=:creator AND name=:name "
|
||||
"AND (:exclude_id IS NULL OR id<>:exclude_id)",
|
||||
{"creator": user_id, "name": name, "exclude_id": exclude_id},
|
||||
)
|
||||
)
|
||||
|
||||
async def dataset_id_conflict(self, dataset_id: str, exclude_id: str) -> bool:
|
||||
return bool(
|
||||
await self.scalar(
|
||||
"SELECT COUNT(*) FROM ai_rag_dataset WHERE dataset_id=:dataset_id AND id<>:exclude_id",
|
||||
{"dataset_id": dataset_id, "exclude_id": exclude_id},
|
||||
)
|
||||
)
|
||||
|
||||
async def insert_dataset(self, values: dict[str, Any]) -> None:
|
||||
await self.execute(
|
||||
"INSERT INTO ai_rag_dataset "
|
||||
"(id,dataset_id,rag_model_id,tenant_id,name,avatar,description,embedding_model,permission,chunk_method,"
|
||||
"parser_config,chunk_count,document_count,token_num,status,creator,created_at,updater,updated_at) VALUES "
|
||||
"(:id,:dataset_id,:rag_model_id,:tenant_id,:name,:avatar,:description,:embedding_model,:permission,"
|
||||
":chunk_method,:parser_config,:chunk_count,:document_count,:token_num,:status,:creator,:created_at,"
|
||||
":updater,:updated_at)",
|
||||
values,
|
||||
)
|
||||
|
||||
async def update_dataset(self, values: dict[str, Any]) -> int:
|
||||
return await self.execute(
|
||||
"UPDATE ai_rag_dataset SET dataset_id=COALESCE(:dataset_id,dataset_id),"
|
||||
"rag_model_id=COALESCE(:rag_model_id,rag_model_id),name=COALESCE(:name,name),"
|
||||
"avatar=COALESCE(:avatar,avatar),description=COALESCE(:description,description),"
|
||||
"embedding_model=COALESCE(:embedding_model,embedding_model),"
|
||||
"permission=COALESCE(:permission,permission),chunk_method=COALESCE(:chunk_method,chunk_method),"
|
||||
"parser_config=COALESCE(:parser_config,parser_config),chunk_count=COALESCE(:chunk_count,chunk_count),"
|
||||
"token_num=COALESCE(:token_num,token_num),status=COALESCE(:status,status),"
|
||||
"creator=COALESCE(:creator,creator),created_at=COALESCE(:created_at,created_at),updater=:updater,"
|
||||
"updated_at=:updated_at WHERE id=:id",
|
||||
values,
|
||||
)
|
||||
|
||||
async def delete_dataset_local(self, row: dict[str, Any]) -> None:
|
||||
await self.execute("DELETE FROM ai_agent_plugin_mapping WHERE plugin_id=:id", {"id": row["id"]})
|
||||
await self.execute("DELETE FROM ai_rag_dataset WHERE id=:id", {"id": row["id"]})
|
||||
|
||||
async def rag_models(self) -> list[dict[str, Any]]:
|
||||
return await self.fetch_all(
|
||||
"SELECT id, model_name, config_json FROM ai_model_config WHERE model_type='RAG' AND is_enabled=1 "
|
||||
"ORDER BY is_default DESC, create_date DESC"
|
||||
)
|
||||
|
||||
async def rag_config(self, model_id: str) -> dict[str, Any]:
|
||||
row = await self.fetch_one("SELECT config_json FROM ai_model_config WHERE id=:id", {"id": model_id})
|
||||
if row is None or row.get("config_json") is None:
|
||||
from app.core.errors import AppError
|
||||
|
||||
raise AppError(10164)
|
||||
raw = row["config_json"]
|
||||
if isinstance(raw, bytes):
|
||||
raw = raw.decode("utf-8")
|
||||
config = dict(raw) if isinstance(raw, dict) else dict(json.loads(str(raw)))
|
||||
config.setdefault("type", "ragflow")
|
||||
return config
|
||||
|
||||
async def documents_page(
|
||||
self,
|
||||
dataset_id: str,
|
||||
*,
|
||||
name: str | None,
|
||||
status: str | None,
|
||||
offset: int,
|
||||
limit: int,
|
||||
) -> tuple[list[dict[str, Any]], int]:
|
||||
where = (
|
||||
"WHERE dataset_id=:dataset_id "
|
||||
"AND (:name IS NULL OR :name='' OR name LIKE CONCAT('%', :name, '%')) "
|
||||
"AND (:status IS NULL OR :status='' OR status=:status)"
|
||||
)
|
||||
params = {"dataset_id": dataset_id, "name": name, "status": status, "offset": offset, "limit": limit}
|
||||
total = int(await self.scalar(f"SELECT COUNT(*) FROM ai_rag_knowledge_document {where}", params) or 0) # noqa: S608
|
||||
rows = await self.fetch_all(
|
||||
f"SELECT * FROM ai_rag_knowledge_document {where} " # noqa: S608
|
||||
"ORDER BY created_at DESC LIMIT :offset, :limit",
|
||||
params,
|
||||
)
|
||||
return rows, total
|
||||
|
||||
async def all_documents(self, dataset_id: str) -> list[dict[str, Any]]:
|
||||
return await self.fetch_all(
|
||||
"SELECT * FROM ai_rag_knowledge_document WHERE dataset_id=:dataset_id", {"dataset_id": dataset_id}
|
||||
)
|
||||
|
||||
async def documents_by_remote_ids(self, dataset_id: str, ids: Sequence[str]) -> list[dict[str, Any]]:
|
||||
if not ids:
|
||||
return []
|
||||
statement = text(
|
||||
"SELECT * FROM ai_rag_knowledge_document WHERE dataset_id=:dataset_id AND document_id IN :ids"
|
||||
).bindparams(bindparam("ids", expanding=True))
|
||||
result = await self.session.execute(statement, {"dataset_id": dataset_id, "ids": list(ids)})
|
||||
return [dict(row) for row in result.mappings().all()]
|
||||
|
||||
async def upsert_document(self, dataset_id: str, remote: dict[str, Any], *, creator: int | None = None) -> bool:
|
||||
document_id = str(remote.get("id") or remote.get("document_id") or "")
|
||||
existing = await self.fetch_one(
|
||||
"SELECT id,created_at FROM ai_rag_knowledge_document WHERE document_id=:id", {"id": document_id}
|
||||
)
|
||||
name = remote.get("name")
|
||||
size = remote.get("size")
|
||||
if size is None:
|
||||
size = remote.get("file_size")
|
||||
meta_fields = remote.get("meta_fields")
|
||||
if meta_fields is None:
|
||||
meta_fields = remote.get("meta")
|
||||
error = remote.get("progress_msg")
|
||||
if error is None:
|
||||
error = remote.get("error")
|
||||
synced_at = _shanghai_now_naive()
|
||||
created_at = remote.get("created_at")
|
||||
if not isinstance(created_at, datetime):
|
||||
created_at = _millis_date(remote.get("create_time"))
|
||||
updated_at = remote.get("updated_at")
|
||||
if not isinstance(updated_at, datetime):
|
||||
updated_at = _millis_date(remote.get("update_time"))
|
||||
values = {
|
||||
"id": existing["id"] if existing else uuid.uuid4().hex,
|
||||
"dataset_id": remote.get("dataset_id") or dataset_id,
|
||||
"document_id": document_id,
|
||||
"name": name,
|
||||
"size": size,
|
||||
"type": _file_type(str(name or "")),
|
||||
"chunk_method": remote.get("chunk_method"),
|
||||
"parser_config": _json_dump(remote.get("parser_config")),
|
||||
"status": _remote_status(remote.get("status")),
|
||||
"run": remote.get("run"),
|
||||
"progress": remote.get("progress"),
|
||||
"thumbnail": remote.get("thumbnail"),
|
||||
"process_duration": remote.get("process_duration"),
|
||||
"meta_fields": _json_dump(meta_fields),
|
||||
"source_type": remote.get("source_type"),
|
||||
"error": error,
|
||||
"chunk_count": remote.get("chunk_count") or 0,
|
||||
"token_count": remote.get("token_count") or 0,
|
||||
"enabled": 1,
|
||||
"creator": creator,
|
||||
"created_at": existing.get("created_at") if existing else (created_at or synced_at),
|
||||
"updated_at": updated_at or synced_at,
|
||||
"synced_at": synced_at,
|
||||
}
|
||||
if existing:
|
||||
await self.execute(
|
||||
"UPDATE ai_rag_knowledge_document SET dataset_id=:dataset_id,document_id=:document_id,name=:name,"
|
||||
"size=:size,type=:type,chunk_method=:chunk_method,parser_config=:parser_config,status=:status,run=:run,"
|
||||
"progress=:progress,thumbnail=:thumbnail,process_duration=:process_duration,meta_fields=:meta_fields,"
|
||||
"source_type=:source_type,error=:error,chunk_count=:chunk_count,token_count=:token_count,enabled=:enabled,"
|
||||
"updated_at=:updated_at,last_sync_at=:synced_at WHERE id=:id",
|
||||
values,
|
||||
)
|
||||
return False
|
||||
await self.execute(
|
||||
"INSERT INTO ai_rag_knowledge_document "
|
||||
"(id,dataset_id,document_id,name,size,type,chunk_method,parser_config,status,run,progress,thumbnail,"
|
||||
"process_duration,meta_fields,source_type,error,chunk_count,token_count,enabled,creator,created_at,updated_at,"
|
||||
"last_sync_at) VALUES (:id,:dataset_id,:document_id,:name,:size,:type,:chunk_method,:parser_config,:status,"
|
||||
":run,:progress,:thumbnail,:process_duration,:meta_fields,:source_type,:error,:chunk_count,:token_count,"
|
||||
":enabled,:creator,COALESCE(:created_at,:synced_at),COALESCE(:updated_at,:synced_at),:synced_at)",
|
||||
values,
|
||||
)
|
||||
return True
|
||||
|
||||
async def update_stats(self, dataset_id: str, docs: int, chunks: int, tokens: int) -> None:
|
||||
await self.execute(
|
||||
"UPDATE ai_rag_dataset SET document_count=document_count+:docs,chunk_count=chunk_count+:chunks,"
|
||||
"token_num=token_num+:tokens,updated_at=:now WHERE dataset_id=:dataset_id",
|
||||
{
|
||||
"dataset_id": dataset_id,
|
||||
"docs": docs,
|
||||
"chunks": chunks,
|
||||
"tokens": tokens,
|
||||
"now": _shanghai_now_naive(),
|
||||
},
|
||||
)
|
||||
|
||||
async def delete_document_shadows(self, dataset_id: str, ids: Sequence[str]) -> int:
|
||||
if not ids:
|
||||
return 0
|
||||
statement = text(
|
||||
"DELETE FROM ai_rag_knowledge_document WHERE dataset_id=:dataset_id AND document_id IN :ids"
|
||||
).bindparams(bindparam("ids", expanding=True))
|
||||
result = await self.session.execute(statement, {"dataset_id": dataset_id, "ids": list(ids)})
|
||||
return int(getattr(result, "rowcount", 0) or 0)
|
||||
|
||||
async def mark_documents_running(self, dataset_id: str, ids: Sequence[str], now: datetime) -> int:
|
||||
if not ids:
|
||||
return 0
|
||||
statement = text(
|
||||
"UPDATE ai_rag_knowledge_document SET run='RUNNING',status='1',updated_at=:now "
|
||||
"WHERE dataset_id=:dataset_id AND document_id IN :ids"
|
||||
).bindparams(bindparam("ids", expanding=True))
|
||||
result = await self.session.execute(statement, {"dataset_id": dataset_id, "ids": list(ids), "now": now})
|
||||
return int(getattr(result, "rowcount", 0) or 0)
|
||||
|
||||
async def mark_document_remote_deleted(self, document_id: str, now: datetime) -> int:
|
||||
return await self.execute(
|
||||
"UPDATE ai_rag_knowledge_document SET run='CANCEL',error=:error,updated_at=:now,last_sync_at=:now "
|
||||
"WHERE document_id=:document_id",
|
||||
{
|
||||
"document_id": document_id,
|
||||
"error": "文档在远程服务中已被删除",
|
||||
"now": now,
|
||||
},
|
||||
)
|
||||
|
||||
async def sync_running_document(
|
||||
self,
|
||||
dataset_id: str,
|
||||
document_id: str,
|
||||
remote: dict[str, Any],
|
||||
now: datetime,
|
||||
) -> int:
|
||||
"""Update exactly the columns touched by Java's status-sync helper."""
|
||||
updated_at = _millis_date(remote.get("update_time")) or now
|
||||
meta_fields = remote.get("meta_fields")
|
||||
assignments = (
|
||||
"status=:status,run=:run,progress=:progress,chunk_count=:chunk_count,token_count=:token_count,"
|
||||
"error=:error,process_duration=:process_duration,thumbnail=:thumbnail,updated_at=:updated_at,"
|
||||
"last_sync_at=:now"
|
||||
)
|
||||
if meta_fields is not None:
|
||||
assignments += ",meta_fields=:meta_fields"
|
||||
return await self.execute(
|
||||
f"UPDATE ai_rag_knowledge_document SET {assignments} " # noqa: S608
|
||||
"WHERE document_id=:document_id AND dataset_id=:dataset_id",
|
||||
{
|
||||
"dataset_id": dataset_id,
|
||||
"document_id": document_id,
|
||||
"status": remote.get("status"),
|
||||
"run": remote.get("run"),
|
||||
"progress": remote.get("progress"),
|
||||
"chunk_count": remote.get("chunk_count"),
|
||||
"token_count": remote.get("token_count"),
|
||||
"error": remote.get("progress_msg") if remote.get("progress_msg") is not None else remote.get("error"),
|
||||
"process_duration": remote.get("process_duration"),
|
||||
"thumbnail": remote.get("thumbnail"),
|
||||
"meta_fields": _json_dump(meta_fields),
|
||||
"updated_at": updated_at,
|
||||
"now": now,
|
||||
},
|
||||
)
|
||||
|
||||
async def running_documents(self) -> list[dict[str, Any]]:
|
||||
return await self.fetch_all("SELECT * FROM ai_rag_knowledge_document WHERE run='RUNNING' AND status='1'")
|
||||
|
||||
|
||||
def _json_dump(value: Any) -> str | None:
|
||||
if value is None:
|
||||
return None
|
||||
if isinstance(value, str):
|
||||
return value
|
||||
return json.dumps(value, ensure_ascii=False, separators=(",", ":"))
|
||||
|
||||
|
||||
def _millis_date(value: Any) -> datetime | None:
|
||||
if value is None:
|
||||
return None
|
||||
try:
|
||||
timezone = ZoneInfo(get_settings().timezone)
|
||||
return datetime.fromtimestamp(float(value) / 1000, timezone).replace(tzinfo=None)
|
||||
except (TypeError, ValueError, OSError):
|
||||
return None
|
||||
|
||||
|
||||
def _file_type(name: str) -> str:
|
||||
last_dot = name.rfind(".")
|
||||
if last_dot <= 0 or last_dot == len(name) - 1:
|
||||
return "unknown"
|
||||
extension = name.rsplit(".", 1)[1].lower()
|
||||
if extension in {"pdf", "doc", "docx", "txt", "md", "mdx"}:
|
||||
return "document"
|
||||
if extension in {"csv", "xls", "xlsx"}:
|
||||
return "spreadsheet"
|
||||
if extension in {"ppt", "pptx"}:
|
||||
return "presentation"
|
||||
return extension
|
||||
|
||||
|
||||
def _remote_status(value: Any) -> str:
|
||||
if value is None or (isinstance(value, str) and not value.strip()):
|
||||
return "1"
|
||||
return str(value)
|
||||
|
||||
|
||||
def _shanghai_now_naive() -> datetime:
|
||||
return datetime.now(ZoneInfo(get_settings().timezone)).replace(tzinfo=None)
|
||||
@@ -0,0 +1,243 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
from collections.abc import Sequence
|
||||
from typing import Any
|
||||
|
||||
from sqlalchemy import bindparam, text
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
|
||||
from app.core.database import Repository
|
||||
|
||||
|
||||
class ModelRepository(Repository):
|
||||
def __init__(self, session: AsyncSession):
|
||||
super().__init__(session)
|
||||
|
||||
async def list_model_names(self, model_type: str, model_name: str | None) -> list[dict[str, Any]]:
|
||||
return await self.fetch_all(
|
||||
"SELECT id, model_name FROM ai_model_config "
|
||||
"WHERE model_type = :model_type AND is_enabled = 1 "
|
||||
"AND (:model_name IS NULL OR :model_name = '' OR model_name LIKE CONCAT('%', :model_name, '%')) "
|
||||
"ORDER BY sort ASC",
|
||||
{"model_type": model_type, "model_name": model_name},
|
||||
)
|
||||
|
||||
async def list_llm_names(self, model_name: str | None) -> list[dict[str, Any]]:
|
||||
return await self.fetch_all(
|
||||
"SELECT id, model_name, config_json FROM ai_model_config "
|
||||
"WHERE model_type = 'llm' AND is_enabled = 1 "
|
||||
"AND (:model_name IS NULL OR :model_name = '' OR model_name LIKE CONCAT('%', :model_name, '%'))",
|
||||
{"model_name": model_name},
|
||||
)
|
||||
|
||||
async def list_providers_by_type(self, model_type: str) -> list[dict[str, Any]]:
|
||||
return await self.fetch_all(
|
||||
"SELECT * FROM ai_model_provider WHERE model_type = :model_type ORDER BY sort ASC",
|
||||
{"model_type": model_type or ""},
|
||||
)
|
||||
|
||||
async def list_providers(
|
||||
self,
|
||||
*,
|
||||
model_type: str | None,
|
||||
name: str | None,
|
||||
offset: int,
|
||||
limit: int,
|
||||
) -> tuple[list[dict[str, Any]], int]:
|
||||
where = (
|
||||
"WHERE (:model_type IS NULL OR :model_type = '' OR model_type = :model_type) "
|
||||
"AND (:name IS NULL OR :name = '' OR name LIKE CONCAT('%', :name, '%') "
|
||||
"OR provider_code LIKE CONCAT('%', :name, '%'))"
|
||||
)
|
||||
params = {"model_type": model_type, "name": name, "offset": offset, "limit": limit}
|
||||
total = int(await self.scalar(f"SELECT COUNT(*) FROM ai_model_provider {where}", params) or 0) # noqa: S608
|
||||
rows = await self.fetch_all(
|
||||
f"SELECT * FROM ai_model_provider {where} " # noqa: S608
|
||||
"ORDER BY model_type ASC, sort ASC LIMIT :offset, :limit",
|
||||
params,
|
||||
)
|
||||
return rows, total
|
||||
|
||||
async def list_model_configs(
|
||||
self,
|
||||
*,
|
||||
model_type: str,
|
||||
model_name: str | None,
|
||||
offset: int,
|
||||
limit: int,
|
||||
) -> tuple[list[dict[str, Any]], int]:
|
||||
where = (
|
||||
"WHERE model_type = :model_type AND "
|
||||
"(:model_name IS NULL OR :model_name = '' OR model_name LIKE CONCAT('%', :model_name, '%'))"
|
||||
)
|
||||
params = {"model_type": model_type, "model_name": model_name, "offset": offset, "limit": limit}
|
||||
total = int(await self.scalar(f"SELECT COUNT(*) FROM ai_model_config {where}", params) or 0) # noqa: S608
|
||||
rows = await self.fetch_all(
|
||||
f"SELECT * FROM ai_model_config {where} " # noqa: S608
|
||||
"ORDER BY is_enabled DESC, sort ASC LIMIT :offset, :limit",
|
||||
params,
|
||||
)
|
||||
return rows, total
|
||||
|
||||
async def get_provider(self, model_type: str, provider_code: str) -> dict[str, Any] | None:
|
||||
return await self.fetch_one(
|
||||
"SELECT * FROM ai_model_provider WHERE model_type = :model_type AND provider_code = :provider_code LIMIT 1",
|
||||
{"model_type": model_type or "", "provider_code": provider_code or ""},
|
||||
)
|
||||
|
||||
async def get_model(self, model_id: str, *, for_update: bool = False) -> dict[str, Any] | None:
|
||||
suffix = " FOR UPDATE" if for_update and self.session.get_bind().dialect.name != "sqlite" else ""
|
||||
return await self.fetch_one(
|
||||
f"SELECT * FROM ai_model_config WHERE id = :id LIMIT 1{suffix}", # noqa: S608
|
||||
{"id": model_id},
|
||||
)
|
||||
|
||||
async def insert_model(self, values: dict[str, Any]) -> None:
|
||||
await self.execute(
|
||||
"INSERT INTO ai_model_config "
|
||||
"(id, model_type, model_code, model_name, is_default, is_enabled, config_json, doc_link, remark, sort) "
|
||||
"VALUES (:id, :model_type, :model_code, :model_name, :is_default, COALESCE(:is_enabled, 0), "
|
||||
":config_json, :doc_link, :remark, COALESCE(:sort, 0))",
|
||||
values,
|
||||
)
|
||||
|
||||
async def update_model(self, values: dict[str, Any]) -> int:
|
||||
return await self.execute(
|
||||
"UPDATE ai_model_config SET model_type=:model_type, model_code=:model_code, "
|
||||
"model_name=COALESCE(:model_name, model_name), is_default=:is_default, "
|
||||
"is_enabled=COALESCE(:is_enabled, is_enabled), config_json=:config_json, doc_link=:doc_link, "
|
||||
"remark=COALESCE(:remark, remark), sort=COALESCE(:sort, sort) WHERE id=:id",
|
||||
values,
|
||||
)
|
||||
|
||||
async def delete_model(self, model_id: str) -> int:
|
||||
return await self.execute("DELETE FROM ai_model_config WHERE id = :id", {"id": model_id})
|
||||
|
||||
async def model_agent_references(self, model_id: str) -> list[str]:
|
||||
rows = await self.fetch_all(
|
||||
"SELECT agent_name FROM ai_agent WHERE vad_model_id=:id OR asr_model_id=:id OR llm_model_id=:id "
|
||||
"OR tts_model_id=:id OR mem_model_id=:id OR vllm_model_id=:id OR intent_model_id=:id",
|
||||
{"id": model_id},
|
||||
)
|
||||
return [str(row.get("agent_name") or "") for row in rows]
|
||||
|
||||
async def intent_reference_count(self, model_id: str) -> int:
|
||||
return int(
|
||||
await self.scalar(
|
||||
"SELECT COUNT(*) FROM ai_model_config WHERE model_type='Intent' AND CAST(config_json AS CHAR) LIKE "
|
||||
"CONCAT('%', :id, '%')",
|
||||
{"id": model_id},
|
||||
)
|
||||
or 0
|
||||
)
|
||||
|
||||
async def set_models_default(self, model_type: str, value: int) -> None:
|
||||
await self.execute(
|
||||
"UPDATE ai_model_config SET is_default=:value WHERE model_type=:model_type",
|
||||
{"value": value, "model_type": model_type},
|
||||
)
|
||||
|
||||
async def set_model_enabled(self, model_id: str, status: int) -> int:
|
||||
return await self.execute(
|
||||
"UPDATE ai_model_config SET is_enabled=:status WHERE id=:id",
|
||||
{"status": status, "id": model_id},
|
||||
)
|
||||
|
||||
async def update_default_template_models(self, model_type: str, model_id: str) -> None:
|
||||
columns = {
|
||||
"ASR": ("asr_model_id",),
|
||||
"VAD": ("vad_model_id",),
|
||||
"LLM": ("llm_model_id",),
|
||||
"TTS": ("tts_model_id", "tts_voice_id"),
|
||||
"VLLM": ("vllm_model_id",),
|
||||
"MEMORY": ("mem_model_id",),
|
||||
"INTENT": ("intent_model_id",),
|
||||
}.get(model_type.upper())
|
||||
if not columns:
|
||||
return
|
||||
if columns == ("tts_model_id", "tts_voice_id"):
|
||||
await self.execute(
|
||||
"UPDATE ai_agent_template SET tts_model_id=:id, tts_voice_id=NULL WHERE sort >= 0",
|
||||
{"id": model_id},
|
||||
)
|
||||
else:
|
||||
column = columns[0]
|
||||
await self.session.execute(
|
||||
text(f"UPDATE ai_agent_template SET {column}=:id WHERE sort >= 0"), # noqa: S608
|
||||
{"id": model_id},
|
||||
)
|
||||
|
||||
async def insert_provider(self, values: dict[str, Any]) -> None:
|
||||
if self.session.get_bind().dialect.name == "sqlite":
|
||||
statement = (
|
||||
"INSERT INTO ai_model_provider "
|
||||
"(id, model_type, provider_code, name, fields, sort, creator, create_date, updater, update_date) "
|
||||
"VALUES (:id, :model_type, :provider_code, :name, :fields, :sort, :creator, :now, :updater, :now)"
|
||||
)
|
||||
else:
|
||||
statement = (
|
||||
"INSERT INTO ai_model_provider "
|
||||
"(id, model_type, provider_code, name, fields, sort, creator, create_date, updater, update_date) "
|
||||
"VALUES (:id, :model_type, :provider_code, :name, CAST(:fields AS JSON), :sort, :creator, :now, "
|
||||
":updater, :now)"
|
||||
)
|
||||
await self.execute(statement, values)
|
||||
|
||||
async def update_provider(self, values: dict[str, Any]) -> int:
|
||||
if self.session.get_bind().dialect.name == "sqlite":
|
||||
statement = (
|
||||
"UPDATE ai_model_provider SET model_type=:model_type, provider_code=:provider_code, name=:name, "
|
||||
"fields=:fields, sort=:sort, updater=:updater, update_date=:now WHERE id=:id"
|
||||
)
|
||||
else:
|
||||
statement = (
|
||||
"UPDATE ai_model_provider SET model_type=:model_type, provider_code=:provider_code, name=:name, "
|
||||
"fields=CAST(:fields AS JSON), sort=:sort, updater=:updater, update_date=:now WHERE id=:id"
|
||||
)
|
||||
return await self.execute(statement, values)
|
||||
|
||||
async def delete_providers(self, ids: Sequence[str]) -> int:
|
||||
if not ids:
|
||||
return 0
|
||||
statement = text("DELETE FROM ai_model_provider WHERE id IN :ids").bindparams(
|
||||
bindparam("ids", expanding=True)
|
||||
)
|
||||
result = await self.session.execute(statement, {"ids": list(ids)})
|
||||
return int(getattr(result, "rowcount", 0) or 0)
|
||||
|
||||
async def list_plugins_for_user(self, user_id: int) -> list[dict[str, Any]]:
|
||||
providers = await self.fetch_all("SELECT * FROM ai_model_provider WHERE model_type='Plugin'")
|
||||
datasets = await self.fetch_all(
|
||||
"SELECT id, name, created_at, updated_at FROM ai_rag_dataset WHERE creator=:creator AND status=1",
|
||||
{"creator": user_id},
|
||||
)
|
||||
providers.extend(
|
||||
{
|
||||
"id": row["id"],
|
||||
"model_type": "Rag",
|
||||
"name": f"[知识库]{row['name']}",
|
||||
"provider_code": "ragflow",
|
||||
"fields": "[]",
|
||||
"sort": 0,
|
||||
"create_date": row.get("created_at"),
|
||||
"update_date": row.get("updated_at"),
|
||||
"creator": 0,
|
||||
"updater": 0,
|
||||
}
|
||||
for row in datasets
|
||||
)
|
||||
return providers
|
||||
|
||||
|
||||
def parse_json_object(value: Any) -> dict[str, Any] | None:
|
||||
if value is None:
|
||||
return None
|
||||
if isinstance(value, dict):
|
||||
return dict(value)
|
||||
if isinstance(value, bytes):
|
||||
value = value.decode("utf-8")
|
||||
if isinstance(value, str):
|
||||
parsed = json.loads(value)
|
||||
return dict(parsed) if isinstance(parsed, dict) else None
|
||||
return None
|
||||
@@ -0,0 +1,153 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from datetime import datetime
|
||||
from typing import Any
|
||||
|
||||
from sqlalchemy import text
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
|
||||
from app.core.database import Repository
|
||||
|
||||
|
||||
class SecurityRepository(Repository):
|
||||
def __init__(self, session: AsyncSession):
|
||||
super().__init__(session)
|
||||
|
||||
async def get_param_value(self, code: str) -> str | None:
|
||||
value = await self.scalar(
|
||||
"SELECT param_value FROM sys_params WHERE param_code = :code LIMIT 1",
|
||||
{"code": code},
|
||||
)
|
||||
return None if value is None else str(value)
|
||||
|
||||
async def get_mobile_area_items(self) -> list[dict[str, Any]]:
|
||||
return await self.fetch_all(
|
||||
"SELECT d.dict_label AS name, d.dict_value AS `key` "
|
||||
"FROM sys_dict_data d "
|
||||
"LEFT JOIN sys_dict_type t ON d.dict_type_id = t.id "
|
||||
"WHERE t.dict_type = :dict_type ORDER BY d.sort ASC",
|
||||
{"dict_type": "MOBILE_AREA"},
|
||||
)
|
||||
|
||||
async def get_user_by_username(self, username: str | None) -> dict[str, Any] | None:
|
||||
return await self.fetch_one(
|
||||
"SELECT id, username, password, super_admin, status, creator, create_date, updater, update_date "
|
||||
"FROM sys_user WHERE username = :username LIMIT 1",
|
||||
{"username": username},
|
||||
)
|
||||
|
||||
async def get_user_by_id(self, user_id: int) -> dict[str, Any] | None:
|
||||
return await self.fetch_one(
|
||||
"SELECT id, username, password, super_admin, status, creator, create_date, updater, update_date "
|
||||
"FROM sys_user WHERE id = :id LIMIT 1",
|
||||
{"id": user_id},
|
||||
)
|
||||
|
||||
async def count_users(self) -> int:
|
||||
return int(await self.scalar("SELECT COUNT(*) FROM sys_user") or 0)
|
||||
|
||||
async def insert_user(
|
||||
self,
|
||||
*,
|
||||
user_id: int,
|
||||
username: str | None,
|
||||
password: str,
|
||||
super_admin: int,
|
||||
now: datetime,
|
||||
) -> None:
|
||||
await self.execute(
|
||||
"INSERT INTO sys_user "
|
||||
"(id, username, password, super_admin, status, creator, create_date, updater, update_date) "
|
||||
"VALUES (:id, :username, :password, :super_admin, 1, NULL, :now, NULL, :now)",
|
||||
{
|
||||
"id": user_id,
|
||||
"username": username,
|
||||
"password": password,
|
||||
"super_admin": super_admin,
|
||||
"now": now,
|
||||
},
|
||||
)
|
||||
|
||||
async def get_token_by_user_id(self, user_id: int, *, for_update: bool = False) -> dict[str, Any] | None:
|
||||
sql = (
|
||||
"SELECT id, user_id, token, expire_date, update_date, create_date "
|
||||
"FROM sys_user_token WHERE user_id = :user_id LIMIT 1 FOR UPDATE"
|
||||
if for_update and self._supports_for_update()
|
||||
else "SELECT id, user_id, token, expire_date, update_date, create_date "
|
||||
"FROM sys_user_token WHERE user_id = :user_id LIMIT 1"
|
||||
)
|
||||
return await self.fetch_one(
|
||||
sql,
|
||||
{"user_id": user_id},
|
||||
)
|
||||
|
||||
async def insert_token(
|
||||
self,
|
||||
*,
|
||||
token_id: int,
|
||||
user_id: int,
|
||||
token: str,
|
||||
now: datetime,
|
||||
expire_date: datetime,
|
||||
) -> None:
|
||||
await self.execute(
|
||||
"INSERT INTO sys_user_token (id, user_id, token, expire_date, update_date, create_date) "
|
||||
"VALUES (:id, :user_id, :token, :expire_date, :now, :now)",
|
||||
{
|
||||
"id": token_id,
|
||||
"user_id": user_id,
|
||||
"token": token,
|
||||
"expire_date": expire_date,
|
||||
"now": now,
|
||||
},
|
||||
)
|
||||
|
||||
async def update_token(self, *, token_id: int, token: str, now: datetime, expire_date: datetime) -> None:
|
||||
await self.execute(
|
||||
"UPDATE sys_user_token SET token = :token, expire_date = :expire_date, update_date = :now "
|
||||
"WHERE id = :id",
|
||||
{"id": token_id, "token": token, "expire_date": expire_date, "now": now},
|
||||
)
|
||||
|
||||
async def update_password(
|
||||
self,
|
||||
user_id: int,
|
||||
password_hash: str,
|
||||
now: datetime,
|
||||
*,
|
||||
preserve_audit_fields: bool = False,
|
||||
) -> int:
|
||||
return await self.execute(
|
||||
"UPDATE sys_user SET password = :password, "
|
||||
"update_date = CASE WHEN :preserve_audit = 1 THEN update_date ELSE :now END WHERE id = :id",
|
||||
{
|
||||
"id": user_id,
|
||||
"password": password_hash,
|
||||
"now": now,
|
||||
"preserve_audit": int(preserve_audit_fields),
|
||||
},
|
||||
)
|
||||
|
||||
async def expire_user_token(self, user_id: int, expire_date: datetime) -> int:
|
||||
return await self.execute(
|
||||
"UPDATE sys_user_token SET expire_date = :expire_date WHERE user_id = :user_id",
|
||||
{"user_id": user_id, "expire_date": expire_date},
|
||||
)
|
||||
|
||||
def _supports_for_update(self) -> bool:
|
||||
bind = self.session.get_bind()
|
||||
return bind.dialect.name != "sqlite"
|
||||
|
||||
|
||||
async def raw_user_token(session: AsyncSession, token: str) -> dict[str, Any] | None:
|
||||
result = await session.execute(
|
||||
text(
|
||||
"SELECT t.id AS token_id, t.user_id, t.token, t.expire_date, "
|
||||
"u.username, u.super_admin, u.status "
|
||||
"FROM sys_user_token t JOIN sys_user u ON u.id = t.user_id "
|
||||
"WHERE t.token = :token LIMIT 1"
|
||||
),
|
||||
{"token": token},
|
||||
)
|
||||
row = result.mappings().first()
|
||||
return dict(row) if row is not None else None
|
||||
@@ -0,0 +1,494 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from datetime import datetime
|
||||
from typing import Any
|
||||
|
||||
from sqlalchemy import bindparam, text
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
|
||||
from app.core.database import Repository
|
||||
|
||||
|
||||
class SysRepository(Repository):
|
||||
def __init__(self, session: AsyncSession):
|
||||
super().__init__(session)
|
||||
|
||||
async def page_users(self, *, mobile: str | None, page: int, limit: int) -> tuple[list[dict[str, Any]], int]:
|
||||
pattern = f"%{mobile}%" if mobile else None
|
||||
params = {"mobile": pattern, "offset": (page - 1) * limit, "limit": limit}
|
||||
total = int(
|
||||
await self.scalar(
|
||||
"SELECT COUNT(*) FROM sys_user WHERE (:mobile IS NULL OR username LIKE :mobile)",
|
||||
params,
|
||||
)
|
||||
or 0
|
||||
)
|
||||
rows = await self.fetch_all(
|
||||
"SELECT u.id, u.username, u.status, u.create_date, "
|
||||
"(SELECT COUNT(*) FROM ai_device d WHERE d.user_id = u.id) AS device_count "
|
||||
"FROM sys_user u WHERE (:mobile IS NULL OR u.username LIKE :mobile) "
|
||||
"ORDER BY u.id ASC LIMIT :limit OFFSET :offset",
|
||||
params,
|
||||
)
|
||||
return rows, total
|
||||
|
||||
async def reset_user_password(
|
||||
self,
|
||||
user_id: int,
|
||||
password_hash: str,
|
||||
updater: int,
|
||||
now: datetime,
|
||||
) -> int:
|
||||
return await self.execute(
|
||||
"UPDATE sys_user SET password = :password, updater = :updater, update_date = :now WHERE id = :id",
|
||||
{"id": user_id, "password": password_hash, "updater": updater, "now": now},
|
||||
)
|
||||
|
||||
async def change_user_status(self, status: int, user_ids: list[int], updater: int, now: datetime) -> int:
|
||||
statement = text(
|
||||
"UPDATE sys_user SET status = :status, updater = :updater, update_date = :now WHERE id IN :ids"
|
||||
).bindparams(bindparam("ids", expanding=True))
|
||||
return await self.execute(
|
||||
statement,
|
||||
{"status": status, "updater": updater, "now": now, "ids": user_ids},
|
||||
)
|
||||
|
||||
async def delete_user_cascade(self, user_id: int) -> None:
|
||||
agent_rows = await self.fetch_all("SELECT id FROM ai_agent WHERE user_id = :user_id", {"user_id": user_id})
|
||||
agent_ids = [str(row["id"]) for row in agent_rows]
|
||||
await self.execute("DELETE FROM sys_user WHERE id = :id", {"id": user_id})
|
||||
await self.execute("DELETE FROM ai_device WHERE user_id = :id", {"id": user_id})
|
||||
for agent_id in agent_ids:
|
||||
audio_rows = await self.fetch_all(
|
||||
"SELECT DISTINCT audio_id FROM ai_agent_chat_history "
|
||||
"WHERE agent_id = :agent_id AND audio_id IS NOT NULL",
|
||||
{"agent_id": agent_id},
|
||||
)
|
||||
audio_ids = [str(row["audio_id"]) for row in audio_rows]
|
||||
if audio_ids:
|
||||
statement = text("DELETE FROM ai_agent_chat_audio WHERE id IN :ids").bindparams(
|
||||
bindparam("ids", expanding=True)
|
||||
)
|
||||
await self.execute(statement, {"ids": audio_ids})
|
||||
for table in (
|
||||
"ai_agent_chat_history",
|
||||
"ai_agent_plugin_mapping",
|
||||
"ai_agent_context_provider",
|
||||
"ai_agent_correct_word_mapping",
|
||||
"ai_agent_tag_relation",
|
||||
"ai_agent_snapshot",
|
||||
):
|
||||
# Table names are a closed list mirroring AgentServiceImpl.deleteAgent.
|
||||
await self.execute(
|
||||
f"DELETE FROM {table} WHERE agent_id = :agent_id", # noqa: S608 - closed table list above
|
||||
{"agent_id": agent_id},
|
||||
)
|
||||
await self.execute("DELETE FROM ai_device WHERE agent_id = :agent_id", {"agent_id": agent_id})
|
||||
await self.execute("DELETE FROM ai_agent WHERE id = :agent_id", {"agent_id": agent_id})
|
||||
|
||||
async def page_devices(
|
||||
self,
|
||||
*,
|
||||
keywords: str | None,
|
||||
page: int,
|
||||
limit: int,
|
||||
) -> tuple[list[dict[str, Any]], int]:
|
||||
pattern = f"%{keywords}%" if keywords else None
|
||||
params = {"keywords": pattern, "offset": (page - 1) * limit, "limit": limit}
|
||||
total = int(
|
||||
await self.scalar(
|
||||
"SELECT COUNT(*) FROM ai_device WHERE (:keywords IS NULL OR alias LIKE :keywords)",
|
||||
params,
|
||||
)
|
||||
or 0
|
||||
)
|
||||
rows = await self.fetch_all(
|
||||
"SELECT d.id, d.user_id, d.mac_address, d.last_connected_at, d.auto_update, d.board, d.alias, "
|
||||
"d.agent_id, d.app_version, d.sort, d.create_date, d.update_date, u.username AS bind_user_name "
|
||||
"FROM ai_device d LEFT JOIN sys_user u ON u.id = d.user_id "
|
||||
"WHERE (:keywords IS NULL OR d.alias LIKE :keywords) "
|
||||
"ORDER BY d.mac_address ASC LIMIT :limit OFFSET :offset",
|
||||
params,
|
||||
)
|
||||
return rows, total
|
||||
|
||||
async def page_params(
|
||||
self,
|
||||
*,
|
||||
param_code: str | None,
|
||||
page: int,
|
||||
limit: int,
|
||||
order_field: str | None,
|
||||
order: str | None,
|
||||
) -> tuple[list[dict[str, Any]], int]:
|
||||
pattern = f"%{param_code}%" if param_code else None
|
||||
params = {"pattern": pattern, "offset": (page - 1) * limit, "limit": limit}
|
||||
where = "param_type = 1 AND (:pattern IS NULL OR param_code LIKE :pattern OR remark LIKE :pattern)"
|
||||
total = int(await self.scalar(f"SELECT COUNT(*) FROM sys_params WHERE {where}", params) or 0) # noqa: S608
|
||||
allowed = {
|
||||
"id": "id",
|
||||
"paramCode": "param_code",
|
||||
"paramValue": "param_value",
|
||||
"valueType": "value_type",
|
||||
"createDate": "create_date",
|
||||
"updateDate": "update_date",
|
||||
}
|
||||
order_column = allowed.get(order_field or "")
|
||||
order_clause = ""
|
||||
if order_column is not None:
|
||||
direction = "ASC" if (order or "").lower() == "asc" else "DESC"
|
||||
order_clause = f" ORDER BY {order_column} {direction}"
|
||||
sql = (
|
||||
"SELECT id, param_code, param_value, value_type, remark, create_date, update_date " # noqa: S608
|
||||
f"FROM sys_params WHERE {where}{order_clause} LIMIT :limit OFFSET :offset"
|
||||
)
|
||||
return await self.fetch_all(sql, params), total # noqa: S608
|
||||
|
||||
async def list_config_params(self) -> list[dict[str, Any]]:
|
||||
return await self.fetch_all(
|
||||
"SELECT id, param_code, param_value, value_type, remark, create_date, update_date "
|
||||
"FROM sys_params WHERE param_type = 1"
|
||||
)
|
||||
|
||||
async def get_param(self, param_id: int) -> dict[str, Any] | None:
|
||||
return await self.fetch_one(
|
||||
"SELECT id, param_code, param_value, value_type, remark, create_date, update_date "
|
||||
"FROM sys_params WHERE id = :id",
|
||||
{"id": param_id},
|
||||
)
|
||||
|
||||
async def get_param_value(self, code: str) -> str | None:
|
||||
value = await self.scalar("SELECT param_value FROM sys_params WHERE param_code = :code", {"code": code})
|
||||
return None if value is None else str(value)
|
||||
|
||||
async def insert_param(
|
||||
self,
|
||||
*,
|
||||
param_id: int,
|
||||
param_code: str,
|
||||
param_value: str,
|
||||
value_type: str,
|
||||
remark: str | None,
|
||||
user_id: int,
|
||||
now: datetime,
|
||||
) -> None:
|
||||
await self.execute(
|
||||
"INSERT INTO sys_params "
|
||||
"(id, param_code, param_value, value_type, param_type, remark, creator, create_date, updater, update_date) "
|
||||
"VALUES (:id, :code, :value, :value_type, 1, :remark, :user_id, :now, :user_id, :now)",
|
||||
{
|
||||
"id": param_id,
|
||||
"code": param_code,
|
||||
"value": param_value,
|
||||
"value_type": value_type,
|
||||
"remark": remark,
|
||||
"user_id": user_id,
|
||||
"now": now,
|
||||
},
|
||||
)
|
||||
|
||||
async def update_param(
|
||||
self,
|
||||
*,
|
||||
param_id: int,
|
||||
param_code: str,
|
||||
param_value: str,
|
||||
value_type: str,
|
||||
remark: str | None,
|
||||
user_id: int,
|
||||
now: datetime,
|
||||
) -> int:
|
||||
return await self.execute(
|
||||
"UPDATE sys_params SET param_code = :code, param_value = :value, value_type = :value_type, "
|
||||
"remark = CASE WHEN :has_remark = 1 THEN :remark ELSE remark END, updater = :user_id, update_date = :now "
|
||||
"WHERE id = :id",
|
||||
{
|
||||
"id": param_id,
|
||||
"code": param_code,
|
||||
"value": param_value,
|
||||
"value_type": value_type,
|
||||
"has_remark": int(remark is not None),
|
||||
"remark": remark,
|
||||
"user_id": user_id,
|
||||
"now": now,
|
||||
},
|
||||
)
|
||||
|
||||
async def update_param_value_by_code(self, code: str, value: str, user_id: int, now: datetime) -> int:
|
||||
return await self.execute(
|
||||
"UPDATE sys_params SET param_value = :value, updater = :user_id, update_date = :now "
|
||||
"WHERE param_code = :code",
|
||||
{"code": code, "value": value, "user_id": user_id, "now": now},
|
||||
)
|
||||
|
||||
async def param_codes_for_ids(self, ids: list[int]) -> list[str]:
|
||||
statement = text("SELECT param_code FROM sys_params WHERE id IN :ids").bindparams(
|
||||
bindparam("ids", expanding=True)
|
||||
)
|
||||
rows = await self.fetch_all(statement, {"ids": ids})
|
||||
return [str(row["param_code"]) for row in rows]
|
||||
|
||||
async def delete_params(self, ids: list[int]) -> int:
|
||||
statement = text("DELETE FROM sys_params WHERE id IN :ids").bindparams(bindparam("ids", expanding=True))
|
||||
return await self.execute(statement, {"ids": ids})
|
||||
|
||||
async def delete_plugin_mapping_by_plugin_id(self, plugin_id: str) -> int:
|
||||
return await self.execute(
|
||||
"DELETE FROM ai_agent_plugin_mapping WHERE plugin_id = :plugin_id",
|
||||
{"plugin_id": plugin_id},
|
||||
)
|
||||
|
||||
async def page_dict_types(
|
||||
self,
|
||||
*,
|
||||
dict_type: str | None,
|
||||
dict_name: str | None,
|
||||
page: int,
|
||||
limit: int,
|
||||
) -> tuple[list[dict[str, Any]], int]:
|
||||
params = {
|
||||
"dict_type": f"%{dict_type}%" if dict_type else None,
|
||||
"dict_name": f"%{dict_name}%" if dict_name else None,
|
||||
"offset": (page - 1) * limit,
|
||||
"limit": limit,
|
||||
}
|
||||
where = (
|
||||
"(:dict_type IS NULL OR t.dict_type LIKE :dict_type) "
|
||||
"AND (:dict_name IS NULL OR t.dict_name LIKE :dict_name)"
|
||||
)
|
||||
total = int(await self.scalar(f"SELECT COUNT(*) FROM sys_dict_type t WHERE {where}", params) or 0) # noqa: S608
|
||||
rows = await self.fetch_all(
|
||||
"SELECT t.id, t.dict_type, t.dict_name, t.remark, t.sort, t.creator, t.create_date, t.updater, " # noqa: S608
|
||||
"t.update_date, creator.username AS creator_name, updater.username AS updater_name "
|
||||
"FROM sys_dict_type t LEFT JOIN sys_user creator ON creator.id = t.creator "
|
||||
"LEFT JOIN sys_user updater ON updater.id = t.updater "
|
||||
f"WHERE {where} ORDER BY t.sort ASC LIMIT :limit OFFSET :offset", # noqa: S608
|
||||
params,
|
||||
)
|
||||
return rows, total
|
||||
|
||||
async def get_dict_type(self, type_id: int) -> dict[str, Any] | None:
|
||||
return await self.fetch_one(
|
||||
"SELECT id, dict_type, dict_name, remark, sort, creator, create_date, updater, update_date "
|
||||
"FROM sys_dict_type WHERE id = :id",
|
||||
{"id": type_id},
|
||||
)
|
||||
|
||||
async def dict_type_exists(self, dict_type: str | None, *, exclude_id: int | None = None) -> bool:
|
||||
count = await self.scalar(
|
||||
"SELECT COUNT(*) FROM sys_dict_type WHERE dict_type = :dict_type "
|
||||
"AND (:exclude_id IS NULL OR id <> :exclude_id)",
|
||||
{"dict_type": dict_type, "exclude_id": exclude_id},
|
||||
)
|
||||
return int(count or 0) > 0
|
||||
|
||||
async def insert_dict_type(
|
||||
self,
|
||||
*,
|
||||
type_id: int,
|
||||
dict_type: str | None,
|
||||
dict_name: str | None,
|
||||
remark: str | None,
|
||||
sort: int | None,
|
||||
user_id: int,
|
||||
now: datetime,
|
||||
) -> None:
|
||||
await self.execute(
|
||||
"INSERT INTO sys_dict_type "
|
||||
"(id, dict_type, dict_name, remark, sort, creator, create_date, updater, update_date) "
|
||||
"VALUES (:id, :dict_type, :dict_name, :remark, :sort, :user_id, :now, :user_id, :now)",
|
||||
{
|
||||
"id": type_id,
|
||||
"dict_type": dict_type,
|
||||
"dict_name": dict_name,
|
||||
"remark": remark,
|
||||
"sort": sort,
|
||||
"user_id": user_id,
|
||||
"now": now,
|
||||
},
|
||||
)
|
||||
|
||||
async def update_dict_type(
|
||||
self,
|
||||
*,
|
||||
type_id: int | None,
|
||||
dict_type: str | None,
|
||||
dict_name: str | None,
|
||||
remark: str | None,
|
||||
sort: int | None,
|
||||
user_id: int,
|
||||
now: datetime,
|
||||
) -> int:
|
||||
return await self.execute(
|
||||
"UPDATE sys_dict_type SET "
|
||||
"dict_type = CASE WHEN :has_dict_type = 1 THEN :dict_type ELSE dict_type END, "
|
||||
"dict_name = CASE WHEN :has_dict_name = 1 THEN :dict_name ELSE dict_name END, "
|
||||
"remark = CASE WHEN :has_remark = 1 THEN :remark ELSE remark END, "
|
||||
"sort = CASE WHEN :has_sort = 1 THEN :sort ELSE sort END, updater = :user_id, update_date = :now "
|
||||
"WHERE id = :id",
|
||||
{
|
||||
"id": type_id,
|
||||
"has_dict_type": int(dict_type is not None),
|
||||
"dict_type": dict_type,
|
||||
"has_dict_name": int(dict_name is not None),
|
||||
"dict_name": dict_name,
|
||||
"has_remark": int(remark is not None),
|
||||
"remark": remark,
|
||||
"has_sort": int(sort is not None),
|
||||
"sort": sort,
|
||||
"user_id": user_id,
|
||||
"now": now,
|
||||
},
|
||||
)
|
||||
|
||||
async def delete_dict_types(self, ids: list[int]) -> None:
|
||||
statement_data = text("DELETE FROM sys_dict_data WHERE dict_type_id IN :ids").bindparams(
|
||||
bindparam("ids", expanding=True)
|
||||
)
|
||||
statement_types = text("DELETE FROM sys_dict_type WHERE id IN :ids").bindparams(
|
||||
bindparam("ids", expanding=True)
|
||||
)
|
||||
await self.execute(statement_data, {"ids": ids})
|
||||
await self.execute(statement_types, {"ids": ids})
|
||||
|
||||
async def page_dict_data(
|
||||
self,
|
||||
*,
|
||||
dict_type_id: int | None,
|
||||
dict_label: str | None,
|
||||
dict_value: str | None,
|
||||
page: int,
|
||||
limit: int,
|
||||
) -> tuple[list[dict[str, Any]], int]:
|
||||
params = {
|
||||
"type_id": dict_type_id,
|
||||
"dict_label": f"%{dict_label}%" if dict_label else None,
|
||||
"dict_value": f"%{dict_value}%" if dict_value else None,
|
||||
"offset": (page - 1) * limit,
|
||||
"limit": limit,
|
||||
}
|
||||
where = (
|
||||
"d.dict_type_id = :type_id AND (:dict_label IS NULL OR d.dict_label LIKE :dict_label) "
|
||||
"AND (:dict_value IS NULL OR d.dict_value LIKE :dict_value)"
|
||||
)
|
||||
total = int(await self.scalar(f"SELECT COUNT(*) FROM sys_dict_data d WHERE {where}", params) or 0) # noqa: S608
|
||||
rows = await self.fetch_all(
|
||||
"SELECT d.id, d.dict_type_id, d.dict_label, d.dict_value, d.remark, d.sort, d.creator, " # noqa: S608
|
||||
"d.create_date, d.updater, d.update_date, creator.username AS creator_name, "
|
||||
"updater.username AS updater_name FROM sys_dict_data d "
|
||||
"LEFT JOIN sys_user creator ON creator.id = d.creator "
|
||||
"LEFT JOIN sys_user updater ON updater.id = d.updater "
|
||||
f"WHERE {where} ORDER BY d.sort ASC LIMIT :limit OFFSET :offset", # noqa: S608
|
||||
params,
|
||||
)
|
||||
return rows, total
|
||||
|
||||
async def get_dict_data(self, data_id: int) -> dict[str, Any] | None:
|
||||
return await self.fetch_one(
|
||||
"SELECT id, dict_type_id, dict_label, dict_value, remark, sort, creator, create_date, updater, update_date "
|
||||
"FROM sys_dict_data WHERE id = :id",
|
||||
{"id": data_id},
|
||||
)
|
||||
|
||||
async def dict_data_label_exists(
|
||||
self,
|
||||
dict_type_id: int | None,
|
||||
compared_label: str | None,
|
||||
*,
|
||||
exclude_id: int | None = None,
|
||||
) -> bool:
|
||||
count = await self.scalar(
|
||||
"SELECT COUNT(*) FROM sys_dict_data WHERE dict_type_id = :type_id AND dict_label = :label "
|
||||
"AND (:exclude_id IS NULL OR id <> :exclude_id)",
|
||||
{"type_id": dict_type_id, "label": compared_label, "exclude_id": exclude_id},
|
||||
)
|
||||
return int(count or 0) > 0
|
||||
|
||||
async def dict_type_code(self, type_id: int | None) -> str | None:
|
||||
value = await self.scalar("SELECT dict_type FROM sys_dict_type WHERE id = :id", {"id": type_id})
|
||||
return None if value is None else str(value)
|
||||
|
||||
async def insert_dict_data(
|
||||
self,
|
||||
*,
|
||||
data_id: int,
|
||||
dict_type_id: int | None,
|
||||
dict_label: str | None,
|
||||
dict_value: str | None,
|
||||
remark: str | None,
|
||||
sort: int | None,
|
||||
user_id: int,
|
||||
now: datetime,
|
||||
) -> None:
|
||||
await self.execute(
|
||||
"INSERT INTO sys_dict_data "
|
||||
"(id, dict_type_id, dict_label, dict_value, remark, sort, creator, create_date, updater, update_date) "
|
||||
"VALUES (:id, :type_id, :label, :value, :remark, :sort, :user_id, :now, :user_id, :now)",
|
||||
{
|
||||
"id": data_id,
|
||||
"type_id": dict_type_id,
|
||||
"label": dict_label,
|
||||
"value": dict_value,
|
||||
"remark": remark,
|
||||
"sort": sort,
|
||||
"user_id": user_id,
|
||||
"now": now,
|
||||
},
|
||||
)
|
||||
|
||||
async def update_dict_data(
|
||||
self,
|
||||
*,
|
||||
data_id: int | None,
|
||||
dict_type_id: int | None,
|
||||
dict_label: str | None,
|
||||
dict_value: str | None,
|
||||
remark: str | None,
|
||||
sort: int | None,
|
||||
user_id: int,
|
||||
now: datetime,
|
||||
) -> int:
|
||||
return await self.execute(
|
||||
"UPDATE sys_dict_data SET "
|
||||
"dict_type_id = CASE WHEN :has_type_id = 1 THEN :type_id ELSE dict_type_id END, "
|
||||
"dict_label = CASE WHEN :has_label = 1 THEN :label ELSE dict_label END, "
|
||||
"dict_value = CASE WHEN :has_value = 1 THEN :value ELSE dict_value END, "
|
||||
"remark = CASE WHEN :has_remark = 1 THEN :remark ELSE remark END, "
|
||||
"sort = CASE WHEN :has_sort = 1 THEN :sort ELSE sort END, updater = :user_id, update_date = :now "
|
||||
"WHERE id = :id",
|
||||
{
|
||||
"id": data_id,
|
||||
"has_type_id": int(dict_type_id is not None),
|
||||
"type_id": dict_type_id,
|
||||
"has_label": int(dict_label is not None),
|
||||
"label": dict_label,
|
||||
"has_value": int(dict_value is not None),
|
||||
"value": dict_value,
|
||||
"has_remark": int(remark is not None),
|
||||
"remark": remark,
|
||||
"has_sort": int(sort is not None),
|
||||
"sort": sort,
|
||||
"user_id": user_id,
|
||||
"now": now,
|
||||
},
|
||||
)
|
||||
|
||||
async def dict_type_codes_for_data_ids(self, ids: list[int]) -> list[str]:
|
||||
statement = text(
|
||||
"SELECT DISTINCT t.dict_type FROM sys_dict_type t JOIN sys_dict_data d ON d.dict_type_id = t.id "
|
||||
"WHERE d.id IN :ids"
|
||||
).bindparams(bindparam("ids", expanding=True))
|
||||
rows = await self.fetch_all(statement, {"ids": ids})
|
||||
return [str(row["dict_type"]) for row in rows]
|
||||
|
||||
async def delete_dict_data(self, ids: list[int]) -> int:
|
||||
statement = text("DELETE FROM sys_dict_data WHERE id IN :ids").bindparams(bindparam("ids", expanding=True))
|
||||
return await self.execute(statement, {"ids": ids})
|
||||
|
||||
async def dict_items(self, dict_type: str) -> list[dict[str, Any]]:
|
||||
return await self.fetch_all(
|
||||
"SELECT d.dict_label AS name, d.dict_value AS `key` FROM sys_dict_data d "
|
||||
"LEFT JOIN sys_dict_type t ON d.dict_type_id = t.id "
|
||||
"WHERE t.dict_type = :dict_type ORDER BY d.sort ASC",
|
||||
{"dict_type": dict_type},
|
||||
)
|
||||
@@ -0,0 +1,71 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from collections.abc import Sequence
|
||||
from typing import Any
|
||||
|
||||
from sqlalchemy import bindparam, text
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
|
||||
from app.core.database import Repository
|
||||
|
||||
|
||||
class TimbreRepository(Repository):
|
||||
def __init__(self, session: AsyncSession):
|
||||
super().__init__(session)
|
||||
|
||||
async def page(
|
||||
self, *, tts_model_id: str, name: str | None, offset: int, limit: int
|
||||
) -> tuple[list[dict[str, Any]], int]:
|
||||
where = (
|
||||
"WHERE tts_model_id=:tts_model_id AND "
|
||||
"(:name IS NULL OR :name='' OR name LIKE CONCAT('%', :name, '%'))"
|
||||
)
|
||||
params = {"tts_model_id": tts_model_id, "name": name, "offset": offset, "limit": limit}
|
||||
total = int(await self.scalar(f"SELECT COUNT(*) FROM ai_tts_voice {where}", params) or 0) # noqa: S608
|
||||
rows = await self.fetch_all(
|
||||
f"SELECT * FROM ai_tts_voice {where} LIMIT :offset, :limit", # noqa: S608
|
||||
params,
|
||||
)
|
||||
return rows, total
|
||||
|
||||
async def insert(self, values: dict[str, Any]) -> None:
|
||||
await self.execute(
|
||||
"INSERT INTO ai_tts_voice "
|
||||
"(id, languages, name, remark, reference_audio, reference_text, sort, tts_model_id, tts_voice, "
|
||||
"voice_demo, creator, create_date) VALUES (:id, :languages, :name, :remark, :reference_audio, "
|
||||
":reference_text, :sort, :tts_model_id, :tts_voice, :voice_demo, :creator, :now)",
|
||||
values,
|
||||
)
|
||||
|
||||
async def update(self, values: dict[str, Any]) -> int:
|
||||
return await self.execute(
|
||||
"UPDATE ai_tts_voice SET languages=:languages, name=:name, remark=COALESCE(:remark, remark), "
|
||||
"reference_audio=COALESCE(:reference_audio, reference_audio), "
|
||||
"reference_text=COALESCE(:reference_text, reference_text), sort=:sort, "
|
||||
"tts_model_id=:tts_model_id, tts_voice=:tts_voice, "
|
||||
"voice_demo=COALESCE(:voice_demo, voice_demo), updater=:updater, "
|
||||
"update_date=:now WHERE id=:id",
|
||||
values,
|
||||
)
|
||||
|
||||
async def delete(self, ids: Sequence[str]) -> int:
|
||||
if not ids:
|
||||
return 0
|
||||
statement = text("DELETE FROM ai_tts_voice WHERE id IN :ids").bindparams(bindparam("ids", expanding=True))
|
||||
result = await self.session.execute(statement, {"ids": list(ids)})
|
||||
return int(getattr(result, "rowcount", 0) or 0)
|
||||
|
||||
async def voices(
|
||||
self, model_id: str, name: str | None, user_id: int
|
||||
) -> tuple[list[dict[str, Any]], list[dict[str, Any]]]:
|
||||
normal = await self.fetch_all(
|
||||
"SELECT id, name, voice_demo, languages FROM ai_tts_voice WHERE tts_model_id=:model_id "
|
||||
"AND (:name IS NULL OR :name='' OR name LIKE CONCAT('%', :name, '%'))",
|
||||
{"model_id": model_id or "", "name": name},
|
||||
)
|
||||
clones = await self.fetch_all(
|
||||
"SELECT id, name, voice_id AS voice_demo, languages FROM ai_voice_clone "
|
||||
"WHERE model_id=:model_id AND user_id=:user_id AND train_status=2",
|
||||
{"model_id": model_id, "user_id": user_id},
|
||||
)
|
||||
return normal, clones
|
||||
@@ -0,0 +1,170 @@
|
||||
from __future__ import annotations
|
||||
|
||||
# Every interpolated SQL fragment below is a module constant or a service-side allowlist.
|
||||
# ruff: noqa: S608
|
||||
from collections.abc import Mapping, Sequence
|
||||
from typing import Any
|
||||
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
|
||||
from app.core.database import Repository
|
||||
|
||||
VOICE_COLUMNS = (
|
||||
"id, name, model_id, voice_id, languages, user_id, voice, train_status, train_error, creator, create_date"
|
||||
)
|
||||
|
||||
|
||||
class VoiceCloneRepository(Repository):
|
||||
def __init__(self, session: AsyncSession):
|
||||
super().__init__(session)
|
||||
|
||||
async def count(self, *, name: str | None, user_id: str | None) -> int:
|
||||
where, params = self._filters(name=name, user_id=user_id)
|
||||
return int(await self.scalar(f"SELECT COUNT(*) FROM ai_voice_clone{where}", params) or 0)
|
||||
|
||||
async def page(
|
||||
self,
|
||||
*,
|
||||
page: int,
|
||||
limit: int,
|
||||
name: str | None,
|
||||
user_id: str | None,
|
||||
order_fields: Sequence[str],
|
||||
ascending: bool,
|
||||
) -> list[dict[str, Any]]:
|
||||
where, params = self._filters(name=name, user_id=user_id)
|
||||
params.update(limit=limit, offset=max(page - 1, 0) * limit)
|
||||
direction = "ASC" if ascending else "DESC"
|
||||
order_by = ", ".join(f"{field} {direction}" for field in order_fields)
|
||||
return await self.fetch_all(
|
||||
f"SELECT {VOICE_COLUMNS} FROM ai_voice_clone{where} "
|
||||
f"ORDER BY {order_by} LIMIT :limit OFFSET :offset",
|
||||
params,
|
||||
)
|
||||
|
||||
async def get(self, voice_id: str | None) -> dict[str, Any] | None:
|
||||
if voice_id is None:
|
||||
return None
|
||||
return await self.fetch_one(
|
||||
f"SELECT {VOICE_COLUMNS} FROM ai_voice_clone WHERE id = :id LIMIT 1",
|
||||
{"id": voice_id},
|
||||
)
|
||||
|
||||
async def list_by_user(self, user_id: int) -> list[dict[str, Any]]:
|
||||
return await self.fetch_all(
|
||||
f"SELECT {VOICE_COLUMNS} FROM ai_voice_clone "
|
||||
"WHERE user_id = :user_id ORDER BY create_date DESC",
|
||||
{"user_id": user_id},
|
||||
)
|
||||
|
||||
async def voice_id_count(self, *, model_id: str, voice_id: str) -> int:
|
||||
return int(
|
||||
await self.scalar(
|
||||
"SELECT COUNT(*) FROM ai_voice_clone WHERE voice_id = :voice_id AND model_id = :model_id",
|
||||
{"model_id": model_id, "voice_id": voice_id},
|
||||
)
|
||||
or 0
|
||||
)
|
||||
|
||||
async def insert_many(self, values: Sequence[Mapping[str, Any]]) -> int:
|
||||
return await self.execute_many(
|
||||
"INSERT INTO ai_voice_clone "
|
||||
"(id, name, model_id, voice_id, languages, user_id, voice, train_status, train_error, creator, "
|
||||
"create_date) VALUES (:id, :name, :model_id, :voice_id, :languages, :user_id, :voice, :train_status, "
|
||||
":train_error, :creator, :create_date)",
|
||||
values,
|
||||
)
|
||||
|
||||
async def delete_many(self, ids: Sequence[str]) -> int:
|
||||
if not ids:
|
||||
return 0
|
||||
placeholders = ", ".join(f":id_{index}" for index in range(len(ids)))
|
||||
params = {f"id_{index}": value for index, value in enumerate(ids)}
|
||||
return await self.execute(f"DELETE FROM ai_voice_clone WHERE id IN ({placeholders})", params)
|
||||
|
||||
async def update_voice(self, voice_id: str, data: bytes) -> int:
|
||||
return await self.execute(
|
||||
"UPDATE ai_voice_clone SET voice = :voice, train_status = 0 WHERE id = :id",
|
||||
{"id": voice_id, "voice": data},
|
||||
)
|
||||
|
||||
async def update_name(self, voice_id: str, name: str) -> int:
|
||||
return await self.execute(
|
||||
"UPDATE ai_voice_clone SET name = :name WHERE id = :id",
|
||||
{"id": voice_id, "name": name},
|
||||
)
|
||||
|
||||
async def update_training(
|
||||
self,
|
||||
voice_id: str,
|
||||
*,
|
||||
train_status: int,
|
||||
train_error: str | None,
|
||||
speaker_id: str | None = None,
|
||||
) -> int:
|
||||
if speaker_id is None:
|
||||
return await self.execute(
|
||||
"UPDATE ai_voice_clone SET train_status = :train_status, train_error = :train_error WHERE id = :id",
|
||||
{"id": voice_id, "train_status": train_status, "train_error": train_error},
|
||||
)
|
||||
return await self.execute(
|
||||
"UPDATE ai_voice_clone SET train_status = :train_status, train_error = :train_error, "
|
||||
"voice_id = :speaker_id WHERE id = :id",
|
||||
{
|
||||
"id": voice_id,
|
||||
"train_status": train_status,
|
||||
"train_error": train_error,
|
||||
"speaker_id": speaker_id,
|
||||
},
|
||||
)
|
||||
|
||||
async def get_model_config(self, model_id: str) -> dict[str, Any] | None:
|
||||
return await self.fetch_one(
|
||||
"SELECT id, model_name, config_json FROM ai_model_config WHERE id = :id LIMIT 1",
|
||||
{"id": model_id},
|
||||
)
|
||||
|
||||
async def get_model_name(self, model_id: str) -> str | None:
|
||||
value = await self.scalar(
|
||||
"SELECT model_name FROM ai_model_config WHERE id = :id LIMIT 1",
|
||||
{"id": model_id},
|
||||
)
|
||||
return None if value is None else str(value)
|
||||
|
||||
async def get_usernames(self, user_ids: Sequence[int]) -> dict[int, str]:
|
||||
if not user_ids:
|
||||
return {}
|
||||
unique_ids = list(dict.fromkeys(user_ids))
|
||||
placeholders = ", ".join(f":user_{index}" for index in range(len(unique_ids)))
|
||||
params = {f"user_{index}": value for index, value in enumerate(unique_ids)}
|
||||
rows = await self.fetch_all(
|
||||
f"SELECT id, username FROM sys_user WHERE id IN ({placeholders})",
|
||||
params,
|
||||
)
|
||||
return {int(row["id"]): str(row["username"]) for row in rows}
|
||||
|
||||
async def get_username(self, user_id: int) -> str | None:
|
||||
value = await self.scalar(
|
||||
"SELECT username FROM sys_user WHERE id = :id LIMIT 1",
|
||||
{"id": user_id},
|
||||
)
|
||||
return None if value is None else str(value)
|
||||
|
||||
async def get_tts_platforms(self) -> list[dict[str, Any]]:
|
||||
return await self.fetch_all(
|
||||
"SELECT id, model_name AS modelName FROM ai_model_config "
|
||||
"WHERE model_type = 'TTS' AND JSON_EXTRACT(config_json, '$.type') = 'huoshan_double_stream'"
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _filters(*, name: str | None, user_id: str | None) -> tuple[str, dict[str, Any]]:
|
||||
clauses: list[str] = []
|
||||
params: dict[str, Any] = {}
|
||||
if user_id is not None and user_id.strip():
|
||||
clauses.append("user_id = :user_id")
|
||||
params["user_id"] = user_id
|
||||
if name is not None and name.strip():
|
||||
clauses.append("(name LIKE :name OR voice_id = :exact_name)")
|
||||
params["name"] = f"%{name}%"
|
||||
params["exact_name"] = name
|
||||
return (" WHERE " + " AND ".join(clauses) if clauses else "", params)
|
||||
@@ -0,0 +1,30 @@
|
||||
"""HTTP routers grouped by the Java business domains."""
|
||||
|
||||
from fastapi import APIRouter
|
||||
|
||||
|
||||
def application_routers() -> list[APIRouter]:
|
||||
"""Return every migrated business router; imports stay explicit for coverage auditing."""
|
||||
from app.routers.agent import router as agent_router
|
||||
from app.routers.config import config_router
|
||||
from app.routers.correctword import correctword_router
|
||||
from app.routers.device import device_router
|
||||
from app.routers.knowledge import knowledge_router
|
||||
from app.routers.model import model_router
|
||||
from app.routers.security import security_router
|
||||
from app.routers.sys import sys_router
|
||||
from app.routers.timbre import timbre_router
|
||||
from app.routers.voiceclone import voiceclone_router
|
||||
|
||||
return [
|
||||
security_router,
|
||||
sys_router,
|
||||
config_router,
|
||||
agent_router,
|
||||
device_router,
|
||||
voiceclone_router,
|
||||
model_router,
|
||||
timbre_router,
|
||||
correctword_router,
|
||||
knowledge_router,
|
||||
]
|
||||
@@ -0,0 +1,386 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Annotated
|
||||
|
||||
from fastapi import APIRouter, BackgroundTasks, Body, Depends, Query, Request
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
from starlette.responses import Response
|
||||
|
||||
from app.core.database import get_db
|
||||
from app.core.errors import ErrorCode
|
||||
from app.core.responses import JavaJSONResponse, error_response, ok
|
||||
from app.core.security import AuthUser, require_normal, require_super_admin
|
||||
from app.schemas.agent import (
|
||||
AgentChatHistoryReport,
|
||||
AgentCreate,
|
||||
AgentMemory,
|
||||
AgentSnapshotPage,
|
||||
AgentSnapshotRestore,
|
||||
AgentTagAssignment,
|
||||
AgentTemplate,
|
||||
AgentUpdate,
|
||||
AgentVoicePrintSave,
|
||||
AgentVoicePrintUpdate,
|
||||
)
|
||||
from app.services.agent import AgentService, run_chat_summary_task
|
||||
|
||||
router = APIRouter(tags=["agent"])
|
||||
DbSession = Annotated[AsyncSession, Depends(get_db)]
|
||||
NormalUser = Annotated[AuthUser, Depends(require_normal)]
|
||||
SuperUser = Annotated[AuthUser, Depends(require_super_admin)]
|
||||
|
||||
|
||||
def _service(session: AsyncSession, user: AuthUser | None, request: Request) -> AgentService:
|
||||
return AgentService(session, user, language=request.headers.get("Accept-Language"))
|
||||
|
||||
|
||||
@router.post("/agent/chat-history/report")
|
||||
async def report_chat_history(report: AgentChatHistoryReport, request: Request, session: DbSession) -> JavaJSONResponse:
|
||||
return ok(await _service(session, None, request).report_chat(report))
|
||||
|
||||
|
||||
@router.post("/agent/chat-history/getDownloadUrl/{agentId}/{sessionId}")
|
||||
async def issue_chat_history_download(
|
||||
agentId: str, sessionId: str, request: Request, session: DbSession, user: NormalUser
|
||||
) -> JavaJSONResponse:
|
||||
service = _service(session, user, request)
|
||||
if not await service.has_agent_permission(agentId):
|
||||
return error_response(request, 10132)
|
||||
return ok(await service.issue_history_token(agentId, sessionId))
|
||||
|
||||
|
||||
@router.get("/agent/chat-history/download/{uuid}/current")
|
||||
async def download_current_chat_history(uuid: str, request: Request, session: DbSession) -> Response:
|
||||
content = await _service(session, None, request).consume_history_download(uuid, previous=False)
|
||||
return Response(
|
||||
content.encode("utf-8"),
|
||||
media_type="text/plain;charset=UTF-8",
|
||||
headers={"Content-Disposition": "attachment;filename=history.txt"},
|
||||
)
|
||||
|
||||
|
||||
@router.get("/agent/chat-history/download/{uuid}/previous")
|
||||
async def download_previous_chat_history(uuid: str, request: Request, session: DbSession) -> Response:
|
||||
content = await _service(session, None, request).consume_history_download(uuid, previous=True)
|
||||
return Response(
|
||||
content.encode("utf-8"),
|
||||
media_type="text/plain;charset=UTF-8",
|
||||
headers={"Content-Disposition": "attachment;filename=history.txt"},
|
||||
)
|
||||
|
||||
|
||||
# Static paths are deliberately registered before /agent/{id}; Starlette resolves in declaration order.
|
||||
@router.get("/agent/template/page")
|
||||
async def template_page(
|
||||
request: Request,
|
||||
session: DbSession,
|
||||
user: SuperUser,
|
||||
page: int = Query(default=1),
|
||||
limit: int = Query(default=10),
|
||||
agentName: str | None = Query(default=None),
|
||||
) -> JavaJSONResponse:
|
||||
return ok(await _service(session, user, request).template_page(page, limit, agentName))
|
||||
|
||||
|
||||
@router.post("/agent/template/batch-remove")
|
||||
async def batch_delete_templates(
|
||||
ids: list[str], request: Request, session: DbSession, user: SuperUser
|
||||
) -> JavaJSONResponse:
|
||||
deleted = await _service(session, user, request).batch_delete_templates(ids)
|
||||
return (
|
||||
ok("批量删除成功") if deleted else error_response(request, ErrorCode.INTERNAL_SERVER_ERROR, "批量删除模板失败")
|
||||
)
|
||||
|
||||
|
||||
@router.get("/agent/template/{id}")
|
||||
async def template_detail(id: str, request: Request, session: DbSession, user: SuperUser) -> JavaJSONResponse:
|
||||
result = await _service(session, user, request).template_detail(id)
|
||||
return ok(result) if result is not None else error_response(request, ErrorCode.INTERNAL_SERVER_ERROR, "模板不存在")
|
||||
|
||||
|
||||
@router.post("/agent/template")
|
||||
async def create_template(
|
||||
template: AgentTemplate, request: Request, session: DbSession, user: SuperUser
|
||||
) -> JavaJSONResponse:
|
||||
return ok(await _service(session, user, request).create_template(template))
|
||||
|
||||
|
||||
@router.put("/agent/template")
|
||||
async def update_template(
|
||||
template: AgentTemplate, request: Request, session: DbSession, user: SuperUser
|
||||
) -> JavaJSONResponse:
|
||||
# MyBatis-Plus raises before returning a boolean when updateById receives
|
||||
# an entity without its @TableId. Keep Java's generic error envelope for
|
||||
# that exact input; an unknown but non-empty id still returns the controller's
|
||||
# explicit "更新模板失败" message below.
|
||||
if template.id is None:
|
||||
return error_response(request, ErrorCode.INTERNAL_SERVER_ERROR)
|
||||
updated = await _service(session, user, request).update_template(template)
|
||||
return ok(template) if updated else error_response(request, ErrorCode.INTERNAL_SERVER_ERROR, "更新模板失败")
|
||||
|
||||
|
||||
@router.delete("/agent/template/{id}")
|
||||
async def delete_template(id: str, request: Request, session: DbSession, user: SuperUser) -> JavaJSONResponse:
|
||||
service = _service(session, user, request)
|
||||
if await service.template_detail(id) is None:
|
||||
return error_response(request, ErrorCode.INTERNAL_SERVER_ERROR, "模板不存在")
|
||||
return (
|
||||
ok("删除模板成功")
|
||||
if await service.delete_template(id)
|
||||
else error_response(request, ErrorCode.INTERNAL_SERVER_ERROR, "删除模板失败")
|
||||
)
|
||||
|
||||
|
||||
@router.post("/agent/voice-print")
|
||||
async def create_voiceprint(
|
||||
dto: AgentVoicePrintSave, request: Request, session: DbSession, user: NormalUser
|
||||
) -> JavaJSONResponse:
|
||||
created = await _service(session, user, request).create_voiceprint(dto)
|
||||
return ok() if created else error_response(request, 10057)
|
||||
|
||||
|
||||
@router.put("/agent/voice-print")
|
||||
async def update_voiceprint(
|
||||
dto: AgentVoicePrintUpdate, request: Request, session: DbSession, user: NormalUser
|
||||
) -> JavaJSONResponse:
|
||||
updated = await _service(session, user, request).update_voiceprint(dto)
|
||||
return ok() if updated else error_response(request, 10058)
|
||||
|
||||
|
||||
@router.delete("/agent/voice-print/{id}")
|
||||
async def delete_voiceprint(id: str, request: Request, session: DbSession, user: NormalUser) -> JavaJSONResponse:
|
||||
deleted = await _service(session, user, request).delete_voiceprint(id)
|
||||
return ok() if deleted else error_response(request, 10059)
|
||||
|
||||
|
||||
@router.get("/agent/voice-print/list/{id}")
|
||||
async def list_voiceprints(id: str, request: Request, session: DbSession, user: NormalUser) -> JavaJSONResponse:
|
||||
return ok(await _service(session, user, request).voiceprint_list(id))
|
||||
|
||||
|
||||
@router.get("/agent/mcp/address/{agentId}")
|
||||
async def mcp_address(agentId: str, request: Request, session: DbSession, user: NormalUser) -> JavaJSONResponse:
|
||||
service = _service(session, user, request)
|
||||
if not await service.has_agent_permission(agentId):
|
||||
return error_response(request, 10200)
|
||||
address = await service.mcp_address(agentId)
|
||||
return ok(address) if address is not None else error_response(request, 10201)
|
||||
|
||||
|
||||
@router.get("/agent/mcp/tools/{agentId}")
|
||||
async def mcp_tools(agentId: str, request: Request, session: DbSession, user: NormalUser) -> JavaJSONResponse:
|
||||
service = _service(session, user, request)
|
||||
if not await service.has_agent_permission(agentId):
|
||||
return error_response(request, 10202)
|
||||
return ok(await service.mcp_tools(agentId))
|
||||
|
||||
|
||||
@router.get("/agent/tag/list")
|
||||
async def all_tags(request: Request, session: DbSession, user: NormalUser) -> JavaJSONResponse:
|
||||
return ok(await _service(session, user, request).all_tags())
|
||||
|
||||
|
||||
@router.post("/agent/tag")
|
||||
async def create_tag(
|
||||
request: Request,
|
||||
session: DbSession,
|
||||
user: NormalUser,
|
||||
params: dict[str, str] = Body(...),
|
||||
) -> JavaJSONResponse:
|
||||
tag_name = params.get("tagName")
|
||||
if tag_name is None or not tag_name.strip():
|
||||
return error_response(request, ErrorCode.INTERNAL_SERVER_ERROR, "标签名称不能为空")
|
||||
return ok(await _service(session, user, request).save_tag(tag_name))
|
||||
|
||||
|
||||
@router.delete("/agent/tag/{id}")
|
||||
async def delete_tag(id: str, request: Request, session: DbSession, user: NormalUser) -> JavaJSONResponse:
|
||||
await _service(session, user, request).delete_tag(id)
|
||||
return ok()
|
||||
|
||||
|
||||
@router.post("/agent/audio/{audioId}")
|
||||
async def issue_audio_token(audioId: str, request: Request, session: DbSession, user: NormalUser) -> JavaJSONResponse:
|
||||
token = await _service(session, user, request).issue_audio_token(audioId)
|
||||
return ok(token) if token is not None else error_response(request, ErrorCode.INTERNAL_SERVER_ERROR, "音频不存在")
|
||||
|
||||
|
||||
@router.get("/agent/play/{uuid}")
|
||||
async def play_agent_audio(uuid: str, request: Request, session: DbSession) -> Response:
|
||||
audio = await _service(session, None, request).consume_audio_token(uuid)
|
||||
if audio is None:
|
||||
return Response(status_code=404)
|
||||
return Response(
|
||||
audio,
|
||||
media_type="application/octet-stream",
|
||||
headers={"Content-Disposition": 'attachment; filename="play.wav"'},
|
||||
)
|
||||
|
||||
|
||||
@router.put("/agent/saveMemory/{macAddress}")
|
||||
async def update_memory(
|
||||
macAddress: str, dto: AgentMemory, request: Request, session: DbSession, user: NormalUser
|
||||
) -> JavaJSONResponse:
|
||||
await _service(session, user, request).update_memory_by_mac(macAddress, dto)
|
||||
return ok()
|
||||
|
||||
|
||||
@router.post("/agent/chat-summary/{sessionId}/save")
|
||||
async def save_chat_summary(
|
||||
sessionId: str, background_tasks: BackgroundTasks, request: Request, session: DbSession
|
||||
) -> JavaJSONResponse:
|
||||
await _service(session, None, request).session_agent(sessionId)
|
||||
background_tasks.add_task(run_chat_summary_task, sessionId)
|
||||
return ok()
|
||||
|
||||
|
||||
@router.post("/agent/chat-title/{sessionId}/generate")
|
||||
async def generate_chat_title(sessionId: str, request: Request, session: DbSession) -> JavaJSONResponse:
|
||||
service = _service(session, None, request)
|
||||
await service.session_agent(sessionId)
|
||||
await service.generate_chat_title(sessionId)
|
||||
return ok()
|
||||
|
||||
|
||||
@router.get("/agent/all")
|
||||
async def admin_agent_list(
|
||||
request: Request,
|
||||
session: DbSession,
|
||||
user: SuperUser,
|
||||
page: int = Query(default=1),
|
||||
limit: int = Query(default=10),
|
||||
orderField: str | None = Query(default=None),
|
||||
order: str | None = Query(default=None),
|
||||
) -> JavaJSONResponse:
|
||||
return ok(await _service(session, user, request).admin_agents(page, limit, orderField, order))
|
||||
|
||||
|
||||
@router.get("/agent/list")
|
||||
async def user_agent_list(
|
||||
request: Request,
|
||||
session: DbSession,
|
||||
user: NormalUser,
|
||||
keyword: str | None = Query(default=None),
|
||||
searchType: str = Query(default="name"),
|
||||
) -> JavaJSONResponse:
|
||||
del searchType # Java accepts the parameter but the consolidated implementation ignores it.
|
||||
return ok(await _service(session, user, request).user_agents(keyword))
|
||||
|
||||
|
||||
@router.get("/agent/template")
|
||||
async def template_list(request: Request, session: DbSession, user: NormalUser) -> JavaJSONResponse:
|
||||
return ok(await _service(session, user, request).templates())
|
||||
|
||||
|
||||
@router.get("/agent/{agentId}/snapshots")
|
||||
async def snapshot_page(
|
||||
agentId: str,
|
||||
request: Request,
|
||||
session: DbSession,
|
||||
user: NormalUser,
|
||||
page: int | None = Query(default=1),
|
||||
limit: int | None = Query(default=10),
|
||||
maxVersionNo: int | None = Query(default=None),
|
||||
) -> JavaJSONResponse:
|
||||
params = AgentSnapshotPage(page=page, limit=limit, max_version_no=maxVersionNo)
|
||||
return ok(await _service(session, user, request).snapshot_page(agentId, params))
|
||||
|
||||
|
||||
@router.get("/agent/{agentId}/snapshots/{snapshotId}")
|
||||
async def snapshot_detail(
|
||||
agentId: str, snapshotId: str, request: Request, session: DbSession, user: NormalUser
|
||||
) -> JavaJSONResponse:
|
||||
return ok(await _service(session, user, request).snapshot_detail(agentId, snapshotId))
|
||||
|
||||
|
||||
@router.post("/agent/{agentId}/snapshots/{snapshotId}/restore")
|
||||
async def restore_snapshot(
|
||||
agentId: str,
|
||||
snapshotId: str,
|
||||
dto: AgentSnapshotRestore,
|
||||
request: Request,
|
||||
session: DbSession,
|
||||
user: NormalUser,
|
||||
) -> JavaJSONResponse:
|
||||
await _service(session, user, request).restore_snapshot(agentId, snapshotId, dto.current_state_token)
|
||||
return ok()
|
||||
|
||||
|
||||
@router.delete("/agent/{agentId}/snapshots/{snapshotId}")
|
||||
async def delete_snapshot(
|
||||
agentId: str, snapshotId: str, request: Request, session: DbSession, user: NormalUser
|
||||
) -> JavaJSONResponse:
|
||||
await _service(session, user, request).delete_snapshot(agentId, snapshotId)
|
||||
return ok()
|
||||
|
||||
|
||||
@router.get("/agent/{id}/sessions")
|
||||
async def agent_sessions(
|
||||
id: str,
|
||||
request: Request,
|
||||
session: DbSession,
|
||||
user: NormalUser,
|
||||
page: str | None = Query(default=None),
|
||||
limit: str | None = Query(default=None),
|
||||
) -> JavaJSONResponse:
|
||||
return ok(await _service(session, user, request).sessions(id, page, limit))
|
||||
|
||||
|
||||
@router.get("/agent/{id}/chat-history/user")
|
||||
async def recent_agent_history(id: str, request: Request, session: DbSession, user: NormalUser) -> JavaJSONResponse:
|
||||
service = _service(session, user, request)
|
||||
if not await service.has_agent_permission(id):
|
||||
return error_response(request, ErrorCode.INTERNAL_SERVER_ERROR, "没有权限查看该智能体的聊天记录")
|
||||
return ok(await service.recent_user_history(id))
|
||||
|
||||
|
||||
@router.get("/agent/{id}/chat-history/audio")
|
||||
async def agent_audio_content(id: str, request: Request, session: DbSession, user: NormalUser) -> JavaJSONResponse:
|
||||
return ok(await _service(session, user, request).audio_content(id))
|
||||
|
||||
|
||||
@router.get("/agent/{id}/chat-history/{sessionId}")
|
||||
async def agent_history(
|
||||
id: str, sessionId: str, request: Request, session: DbSession, user: NormalUser
|
||||
) -> JavaJSONResponse:
|
||||
service = _service(session, user, request)
|
||||
if not await service.has_agent_permission(id):
|
||||
return error_response(request, ErrorCode.INTERNAL_SERVER_ERROR, "没有权限查看该智能体的聊天记录")
|
||||
return ok(await service.history(id, sessionId))
|
||||
|
||||
|
||||
@router.get("/agent/{id}/tags")
|
||||
async def agent_tags(id: str, request: Request, session: DbSession, user: NormalUser) -> JavaJSONResponse:
|
||||
return ok(await _service(session, user, request).agent_tags(id))
|
||||
|
||||
|
||||
@router.put("/agent/{id}/tags")
|
||||
async def save_agent_tags(
|
||||
id: str, dto: AgentTagAssignment, request: Request, session: DbSession, user: NormalUser
|
||||
) -> JavaJSONResponse:
|
||||
await _service(session, user, request).save_agent_tags(id, dto.tag_ids, dto.tag_names)
|
||||
return ok()
|
||||
|
||||
|
||||
@router.post("/agent")
|
||||
async def create_agent(dto: AgentCreate, request: Request, session: DbSession, user: NormalUser) -> JavaJSONResponse:
|
||||
return ok(await _service(session, user, request).create_agent(dto))
|
||||
|
||||
|
||||
@router.put("/agent/{id}")
|
||||
async def update_agent(
|
||||
id: str, dto: AgentUpdate, request: Request, session: DbSession, user: NormalUser
|
||||
) -> JavaJSONResponse:
|
||||
await _service(session, user, request).update_agent(id, dto)
|
||||
return ok()
|
||||
|
||||
|
||||
@router.delete("/agent/{id}")
|
||||
async def delete_agent(id: str, request: Request, session: DbSession, user: NormalUser) -> JavaJSONResponse:
|
||||
await _service(session, user, request).delete_agent(id)
|
||||
return ok()
|
||||
|
||||
|
||||
@router.get("/agent/{id}")
|
||||
async def agent_detail(id: str, request: Request, session: DbSession, user: NormalUser) -> JavaJSONResponse:
|
||||
return ok(await _service(session, user, request).agent_detail(id))
|
||||
@@ -0,0 +1,34 @@
|
||||
# ruff: noqa: B008
|
||||
# FastAPI evaluates dependency marker defaults intentionally when registering routes.
|
||||
from __future__ import annotations
|
||||
|
||||
from fastapi import APIRouter, Depends
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
|
||||
from app.core.database import get_db
|
||||
from app.core.responses import JavaJSONResponse, ok
|
||||
from app.core.serialization import preserve_java_map_keys
|
||||
from app.repositories.config import ConfigRepository
|
||||
from app.schemas.config import AgentModelsRequest, CorrectWordsRequest
|
||||
from app.services.config import ConfigService
|
||||
|
||||
config_router = APIRouter()
|
||||
|
||||
|
||||
def _service(session: AsyncSession) -> ConfigService:
|
||||
return ConfigService(ConfigRepository(session))
|
||||
|
||||
|
||||
@config_router.post("/config/server-base")
|
||||
async def server_base(session: AsyncSession = Depends(get_db)) -> JavaJSONResponse:
|
||||
return ok(preserve_java_map_keys(await _service(session).get_config(use_cache=True)))
|
||||
|
||||
|
||||
@config_router.post("/config/agent-models")
|
||||
async def agent_models(dto: AgentModelsRequest, session: AsyncSession = Depends(get_db)) -> JavaJSONResponse:
|
||||
return ok(preserve_java_map_keys(await _service(session).get_agent_models(dto.mac_address, dto.selected_module)))
|
||||
|
||||
|
||||
@config_router.post("/config/correct-words")
|
||||
async def correct_words(dto: CorrectWordsRequest, session: AsyncSession = Depends(get_db)) -> JavaJSONResponse:
|
||||
return ok(await _service(session).get_correct_words(dto.mac_address))
|
||||
@@ -0,0 +1,97 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from urllib.parse import quote
|
||||
|
||||
from fastapi import APIRouter, Depends, Request
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
from starlette.responses import Response
|
||||
|
||||
from app.core.database import get_db
|
||||
from app.core.responses import JavaJSONResponse, ok
|
||||
from app.core.security import require_normal
|
||||
from app.repositories.correctword import CorrectWordRepository
|
||||
from app.schemas.correctword import CorrectWordFileBody
|
||||
from app.services.correctword import CorrectWordService
|
||||
|
||||
correctword_router = APIRouter()
|
||||
|
||||
|
||||
def _java_urlencode(value: str) -> str:
|
||||
# java.net.URLEncoder leaves alphanumerics plus .-*_ unescaped, encodes
|
||||
# spaces as '+', and encodes '~'. The controller then replaces '+' with
|
||||
# '%20'. urllib always leaves '~', so handle that final difference here.
|
||||
return quote(value, safe="*.-_").replace("~", "%7E")
|
||||
|
||||
|
||||
def _service(session: AsyncSession) -> CorrectWordService:
|
||||
return CorrectWordService(CorrectWordRepository(session))
|
||||
|
||||
|
||||
@correctword_router.post("/correct-word/file")
|
||||
async def create_file(
|
||||
body: CorrectWordFileBody, request: Request, session: AsyncSession = Depends(get_db)
|
||||
) -> JavaJSONResponse:
|
||||
return ok(await _service(session).create(body, require_normal(request)))
|
||||
|
||||
|
||||
@correctword_router.put("/correct-word/file/{file_id}")
|
||||
async def update_file(
|
||||
file_id: str,
|
||||
body: CorrectWordFileBody,
|
||||
request: Request,
|
||||
session: AsyncSession = Depends(get_db),
|
||||
) -> JavaJSONResponse:
|
||||
await _service(session).update(file_id, body, require_normal(request))
|
||||
return ok()
|
||||
|
||||
|
||||
@correctword_router.get("/correct-word/file/list")
|
||||
async def list_files(
|
||||
request: Request,
|
||||
page: str | None = None,
|
||||
limit: str | None = None,
|
||||
session: AsyncSession = Depends(get_db),
|
||||
) -> JavaJSONResponse:
|
||||
return ok(await _service(session).page(require_normal(request), page, limit))
|
||||
|
||||
|
||||
@correctword_router.get("/correct-word/file/select")
|
||||
async def select_files(request: Request, session: AsyncSession = Depends(get_db)) -> JavaJSONResponse:
|
||||
return ok(await _service(session).all(require_normal(request)))
|
||||
|
||||
|
||||
@correctword_router.get("/correct-word/file/download/{file_id}")
|
||||
async def download_file(
|
||||
file_id: str, request: Request, session: AsyncSession = Depends(get_db)
|
||||
) -> Response:
|
||||
require_normal(request)
|
||||
item = await _service(session).get(file_id)
|
||||
if item is None or not item["content"]:
|
||||
return Response(status_code=404)
|
||||
body = "\n".join(item["content"]).encode("utf-8")
|
||||
file_name = str(item["fileName"])
|
||||
ascii_name = "".join(character if ord(character) < 128 else "_" for character in file_name)
|
||||
disposition = f"attachment; filename=\"{ascii_name}\"; filename*=UTF-8''{_java_urlencode(file_name)}"
|
||||
return Response(
|
||||
body,
|
||||
media_type="application/octet-stream",
|
||||
headers={"Content-Disposition": disposition, "Content-Length": str(len(body))},
|
||||
)
|
||||
|
||||
|
||||
@correctword_router.delete("/correct-word/file/{file_id}")
|
||||
async def delete_file(
|
||||
file_id: str, request: Request, session: AsyncSession = Depends(get_db)
|
||||
) -> JavaJSONResponse:
|
||||
require_normal(request)
|
||||
await _service(session).delete([file_id])
|
||||
return ok()
|
||||
|
||||
|
||||
@correctword_router.post("/correct-word/file/batch-delete")
|
||||
async def batch_delete_files(
|
||||
file_ids: list[str], request: Request, session: AsyncSession = Depends(get_db)
|
||||
) -> JavaJSONResponse:
|
||||
require_normal(request)
|
||||
await _service(session).delete(file_ids)
|
||||
return ok()
|
||||
@@ -0,0 +1,505 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
from typing import Annotated, Any
|
||||
|
||||
from fastapi import APIRouter, BackgroundTasks, Depends, File, Header, Query, Request, UploadFile
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
from starlette.responses import Response
|
||||
|
||||
from app.core.database import get_db
|
||||
from app.core.i18n import resolve_language
|
||||
from app.core.responses import JavaJSONResponse, envelope, error_response, ok
|
||||
from app.core.security import require_normal, require_super_admin
|
||||
from app.schemas.device import (
|
||||
DeviceAddressBookAliasRequest,
|
||||
DeviceAddressBookPermissionRequest,
|
||||
DeviceManualAddRequest,
|
||||
DeviceRegisterRequest,
|
||||
DeviceReportRequest,
|
||||
DeviceToolCallRequest,
|
||||
DeviceUnbindRequest,
|
||||
DeviceUpdateRequest,
|
||||
OtaRecord,
|
||||
)
|
||||
from app.services.device import MAC_PATTERN, DeviceService, is_blank
|
||||
|
||||
device_router = APIRouter()
|
||||
SessionDep = Annotated[AsyncSession, Depends(get_db)]
|
||||
FirmwareUpload = Annotated[UploadFile, File()]
|
||||
CallerMacQuery = Annotated[str, Query(alias="callerMac")]
|
||||
DeviceIdHeader = Annotated[str | None, Header(alias="Device-Id")]
|
||||
ClientIdHeader = Annotated[str | None, Header(alias="Client-Id")]
|
||||
|
||||
|
||||
def _query_map(request: Request) -> dict[str, Any]:
|
||||
result: dict[str, Any] = {}
|
||||
for key, value in request.query_params.multi_items():
|
||||
if key in result:
|
||||
previous = result[key]
|
||||
result[key] = [*previous, value] if isinstance(previous, list) else [previous, value]
|
||||
else:
|
||||
result[key] = value
|
||||
return result
|
||||
|
||||
|
||||
def _raw_ota(payload: dict[str, Any]) -> Response:
|
||||
body = json.dumps(payload, ensure_ascii=False, separators=(",", ":")).encode("utf-8")
|
||||
return Response(
|
||||
body,
|
||||
status_code=200,
|
||||
media_type="application/json",
|
||||
headers={"Content-Length": str(len(body))},
|
||||
)
|
||||
|
||||
|
||||
@device_router.post("/device/bind/{agent_id}/{device_code}")
|
||||
async def bind_device(
|
||||
agent_id: str,
|
||||
device_code: str,
|
||||
request: Request,
|
||||
session: SessionDep,
|
||||
) -> JavaJSONResponse:
|
||||
user = require_normal(request)
|
||||
await DeviceService(session).activate_bound_device(agent_id=agent_id, activation_code=device_code, user=user)
|
||||
return ok()
|
||||
|
||||
|
||||
@device_router.post("/device/register")
|
||||
async def register_device(
|
||||
body: DeviceRegisterRequest,
|
||||
request: Request,
|
||||
session: SessionDep,
|
||||
) -> JavaJSONResponse:
|
||||
require_normal(request)
|
||||
if is_blank(body.mac_address):
|
||||
return error_response(request, 10175)
|
||||
return ok(await DeviceService(session).register_device(body.mac_address or ""))
|
||||
|
||||
|
||||
@device_router.get("/device/bind/{agent_id}")
|
||||
async def get_bound_devices(
|
||||
agent_id: str,
|
||||
request: Request,
|
||||
session: SessionDep,
|
||||
) -> JavaJSONResponse:
|
||||
user = require_normal(request)
|
||||
return ok(await DeviceService(session).list_user_devices(user.id, agent_id))
|
||||
|
||||
|
||||
@device_router.post("/device/bind/{agent_id}")
|
||||
async def device_online(
|
||||
agent_id: str,
|
||||
request: Request,
|
||||
session: SessionDep,
|
||||
) -> JavaJSONResponse:
|
||||
user = require_normal(request)
|
||||
await request.body()
|
||||
try:
|
||||
return ok(await DeviceService(session).get_online_data(agent_id, user))
|
||||
except Exception as exc:
|
||||
return error_response(request, 500, f"转发请求失败: {exc}")
|
||||
|
||||
|
||||
@device_router.post("/device/unbind")
|
||||
async def unbind_device(
|
||||
body: DeviceUnbindRequest,
|
||||
request: Request,
|
||||
session: SessionDep,
|
||||
) -> JavaJSONResponse:
|
||||
user = require_normal(request)
|
||||
# DeviceController does not apply @Valid to DeviceUnBindDTO. An empty
|
||||
# object reaches the service with a null id and is a successful no-op.
|
||||
await DeviceService(session).unbind(user_id=user.id, device_id=body.device_id or "")
|
||||
return ok()
|
||||
|
||||
|
||||
@device_router.put("/device/update/{device_id}")
|
||||
async def update_device(
|
||||
device_id: str,
|
||||
body: DeviceUpdateRequest,
|
||||
request: Request,
|
||||
session: SessionDep,
|
||||
) -> JavaJSONResponse:
|
||||
user = require_normal(request)
|
||||
validation = _validate_device_update(body, request.headers.get("Accept-Language"))
|
||||
if validation is not None:
|
||||
return error_response(request, 10034, validation)
|
||||
if not await DeviceService(session).update_device(device_id=device_id, request=body, user=user):
|
||||
return error_response(request, 500, "设备不存在")
|
||||
return ok()
|
||||
|
||||
|
||||
@device_router.put("/user/configDevice/{device_id}")
|
||||
async def configure_device(
|
||||
device_id: str,
|
||||
body: DeviceUpdateRequest,
|
||||
request: Request,
|
||||
session: SessionDep,
|
||||
) -> JavaJSONResponse:
|
||||
user = require_normal(request)
|
||||
validation = _validate_device_update(body, request.headers.get("Accept-Language"))
|
||||
if validation is not None:
|
||||
return error_response(request, 10034, validation)
|
||||
if not await DeviceService(session).update_device(device_id=device_id, request=body, user=user):
|
||||
return error_response(request, 500, "设备不存在")
|
||||
return ok()
|
||||
|
||||
|
||||
@device_router.post("/device/manual-add")
|
||||
async def manual_add_device(
|
||||
body: DeviceManualAddRequest,
|
||||
request: Request,
|
||||
session: SessionDep,
|
||||
) -> JavaJSONResponse:
|
||||
user = require_normal(request)
|
||||
await DeviceService(session).manual_add(request=body, user=user)
|
||||
return ok()
|
||||
|
||||
|
||||
@device_router.post("/device/tools/list/{device_id}")
|
||||
async def list_device_tools(
|
||||
device_id: str,
|
||||
request: Request,
|
||||
session: SessionDep,
|
||||
) -> JavaJSONResponse:
|
||||
user = require_normal(request)
|
||||
tools = await DeviceService(session).get_tools(device_id=device_id, user=user)
|
||||
if tools is None:
|
||||
return error_response(request, 10194)
|
||||
return ok(tools)
|
||||
|
||||
|
||||
@device_router.post("/device/tools/call/{device_id}")
|
||||
async def call_device_tool(
|
||||
device_id: str,
|
||||
body: DeviceToolCallRequest,
|
||||
request: Request,
|
||||
session: SessionDep,
|
||||
) -> JavaJSONResponse:
|
||||
user = require_normal(request)
|
||||
if is_blank(body.name):
|
||||
return error_response(request, 10034, "工具名称不能为空")
|
||||
result = await DeviceService(session).call_tool(
|
||||
device_id=device_id,
|
||||
tool_name=body.name or "",
|
||||
arguments=body.arguments,
|
||||
user=user,
|
||||
)
|
||||
if result is None:
|
||||
return error_response(request, 10194)
|
||||
return JavaJSONResponse(envelope(result, msg="Tools called successfully"))
|
||||
|
||||
|
||||
# Static address-book paths deliberately precede /address-book/{mac_address}.
|
||||
@device_router.get("/device/address-book/call")
|
||||
async def call_address_book(
|
||||
request: Request,
|
||||
session: SessionDep,
|
||||
caller_mac: CallerMacQuery,
|
||||
nickname: str,
|
||||
answer: bool = False,
|
||||
) -> JavaJSONResponse:
|
||||
result = await DeviceService(session).call_by_nickname(
|
||||
caller_mac=caller_mac,
|
||||
nickname=nickname,
|
||||
answer=answer,
|
||||
)
|
||||
return ok(result)
|
||||
|
||||
|
||||
@device_router.get("/device/address-book/lookup")
|
||||
async def lookup_address_book(
|
||||
request: Request,
|
||||
session: SessionDep,
|
||||
caller_mac: CallerMacQuery,
|
||||
nickname: str,
|
||||
) -> JavaJSONResponse:
|
||||
result = await DeviceService(session).lookup_address_book(caller_mac=caller_mac, nickname=nickname)
|
||||
if result is None:
|
||||
return error_response(request, 500, "未找到对应设备")
|
||||
return ok(result)
|
||||
|
||||
|
||||
@device_router.put("/device/address-book/alias")
|
||||
async def update_address_alias(
|
||||
body: DeviceAddressBookAliasRequest,
|
||||
request: Request,
|
||||
session: SessionDep,
|
||||
) -> JavaJSONResponse:
|
||||
user = require_normal(request)
|
||||
if is_blank(body.target_mac):
|
||||
return error_response(request, 10034, "目标MAC地址不能为空")
|
||||
if is_blank(body.mac_address):
|
||||
return error_response(request, 10034, "MAC地址不能为空")
|
||||
service = DeviceService(session)
|
||||
caller = await service.repository.get_device_by_mac(body.mac_address or "")
|
||||
if caller is None or int(caller.get("user_id") or -1) != user.id:
|
||||
return error_response(request, 500, "无权限操作该设备")
|
||||
await service.save_address_book(
|
||||
mac_address=body.mac_address or "",
|
||||
target_mac=body.target_mac or "",
|
||||
alias=body.alias,
|
||||
has_permission=None,
|
||||
actor=user.id,
|
||||
)
|
||||
return ok()
|
||||
|
||||
|
||||
@device_router.put("/device/address-book/permission")
|
||||
async def update_address_permission(
|
||||
body: DeviceAddressBookPermissionRequest,
|
||||
request: Request,
|
||||
session: SessionDep,
|
||||
) -> JavaJSONResponse:
|
||||
user = require_normal(request)
|
||||
if is_blank(body.mac_address):
|
||||
return error_response(request, 10034, "MAC地址不能为空")
|
||||
if is_blank(body.target_mac):
|
||||
return error_response(request, 10034, "目标MAC地址不能为空")
|
||||
service = DeviceService(session)
|
||||
caller = await service.repository.get_device_by_mac(body.mac_address or "")
|
||||
if caller is None or int(caller.get("user_id") or -1) != user.id:
|
||||
return error_response(request, 500, "无权限操作该设备")
|
||||
await service.save_address_book(
|
||||
mac_address=body.mac_address or "",
|
||||
target_mac=body.target_mac or "",
|
||||
alias=None,
|
||||
has_permission=body.has_permission,
|
||||
actor=user.id,
|
||||
)
|
||||
return ok()
|
||||
|
||||
|
||||
@device_router.get("/device/address-book/{mac_address}")
|
||||
async def get_address_book(
|
||||
mac_address: str,
|
||||
request: Request,
|
||||
session: SessionDep,
|
||||
) -> JavaJSONResponse:
|
||||
require_normal(request)
|
||||
return ok(await DeviceService(session).address_book(mac_address))
|
||||
|
||||
|
||||
@device_router.post("/ota/")
|
||||
async def check_ota_version(
|
||||
report: DeviceReportRequest,
|
||||
request: Request,
|
||||
session: SessionDep,
|
||||
background_tasks: BackgroundTasks,
|
||||
device_id: DeviceIdHeader = None,
|
||||
client_id: ClientIdHeader = None,
|
||||
) -> Response:
|
||||
if is_blank(device_id):
|
||||
# Java's required @RequestHeader fails before the controller's blank
|
||||
# guard and is translated by its global handler into this envelope.
|
||||
return error_response(request, 500)
|
||||
if MAC_PATTERN.fullmatch(device_id or "") is None:
|
||||
return _raw_ota({"error": "Invalid device ID"})
|
||||
selected_client = device_id if is_blank(client_id) else client_id
|
||||
client_ip = request.client.host if request.client is not None else "unknown"
|
||||
service = DeviceService(session)
|
||||
|
||||
def defer_connection_update(device: str, agent: str | None, version: str | None) -> None:
|
||||
background_tasks.add_task(
|
||||
DeviceService.persist_connection_update_background,
|
||||
device,
|
||||
agent,
|
||||
version,
|
||||
)
|
||||
|
||||
payload = await service.check_ota(
|
||||
device_id=device_id or "",
|
||||
client_id=selected_client or device_id or "",
|
||||
report=report,
|
||||
request_url=str(request.url),
|
||||
client_ip=client_ip,
|
||||
defer_connection_update=defer_connection_update,
|
||||
)
|
||||
return _raw_ota(payload)
|
||||
|
||||
|
||||
@device_router.post("/ota/activate")
|
||||
async def activate_ota_device(
|
||||
request: Request,
|
||||
session: SessionDep,
|
||||
device_id: DeviceIdHeader = None,
|
||||
client_id: ClientIdHeader = None,
|
||||
) -> Response:
|
||||
del client_id
|
||||
if is_blank(device_id):
|
||||
return error_response(request, 500)
|
||||
if await DeviceService(session).repository.get_device_by_mac(device_id or "") is None:
|
||||
return Response(status_code=202)
|
||||
return Response("success", media_type="text/plain;charset=UTF-8")
|
||||
|
||||
|
||||
@device_router.get("/ota/")
|
||||
async def ota_health(session: SessionDep) -> Response:
|
||||
return Response(
|
||||
await DeviceService(session).ota_health_text(),
|
||||
media_type="text/plain;charset=UTF-8",
|
||||
)
|
||||
|
||||
|
||||
# Static otaMag paths deliberately precede /otaMag/{id}.
|
||||
@device_router.get("/otaMag/getDownloadUrl/{ota_id}")
|
||||
async def get_ota_download_url(
|
||||
ota_id: str,
|
||||
request: Request,
|
||||
session: SessionDep,
|
||||
) -> JavaJSONResponse:
|
||||
require_super_admin(request)
|
||||
return ok(await DeviceService(session).create_ota_download_id(ota_id))
|
||||
|
||||
|
||||
@device_router.get("/otaMag/download/{download_id}")
|
||||
async def download_ota(download_id: str, session: SessionDep) -> Response:
|
||||
resolved = await DeviceService(session).resolve_ota_download(download_id)
|
||||
if resolved is None:
|
||||
return Response(status_code=404)
|
||||
path, filename = resolved
|
||||
try:
|
||||
content = path.read_bytes()
|
||||
except OSError:
|
||||
return Response(status_code=500)
|
||||
return Response(
|
||||
content,
|
||||
media_type="application/octet-stream",
|
||||
headers={
|
||||
"Content-Disposition": f'attachment; filename="{filename}"',
|
||||
"Content-Length": str(len(content)),
|
||||
},
|
||||
)
|
||||
|
||||
|
||||
@device_router.post("/otaMag/upload")
|
||||
async def upload_firmware(
|
||||
request: Request,
|
||||
file: FirmwareUpload,
|
||||
session: SessionDep,
|
||||
) -> JavaJSONResponse:
|
||||
require_super_admin(request)
|
||||
service = DeviceService(session)
|
||||
try:
|
||||
content = await file.read()
|
||||
return ok(await service.save_firmware_file(filename=file.filename, content=content))
|
||||
except ValueError as exc:
|
||||
return error_response(request, 500, str(exc))
|
||||
except OSError as exc:
|
||||
return error_response(request, 500, f"文件上传失败:{exc}")
|
||||
|
||||
|
||||
@device_router.post("/otaMag/uploadAssetsBin")
|
||||
async def upload_assets_firmware(
|
||||
request: Request,
|
||||
file: FirmwareUpload,
|
||||
session: SessionDep,
|
||||
) -> JavaJSONResponse:
|
||||
user = require_normal(request)
|
||||
service = DeviceService(session)
|
||||
try:
|
||||
content = await file.read()
|
||||
return ok(await service.save_assets_file(filename=file.filename, content=content, user=user))
|
||||
except ValueError as exc:
|
||||
return error_response(request, 500, str(exc))
|
||||
except OSError as exc:
|
||||
return error_response(request, 500, f"文件上传失败:{exc}")
|
||||
|
||||
|
||||
@device_router.get("/otaMag")
|
||||
async def page_ota(
|
||||
request: Request,
|
||||
session: SessionDep,
|
||||
) -> JavaJSONResponse:
|
||||
require_super_admin(request)
|
||||
return ok(await DeviceService(session).ota_page(_query_map(request)))
|
||||
|
||||
|
||||
@device_router.get("/otaMag/{ota_id}")
|
||||
async def get_ota(
|
||||
ota_id: str,
|
||||
request: Request,
|
||||
session: SessionDep,
|
||||
) -> JavaJSONResponse:
|
||||
require_super_admin(request)
|
||||
return ok(await DeviceService(session).get_ota_record(ota_id))
|
||||
|
||||
|
||||
@device_router.post("/otaMag")
|
||||
async def save_ota(
|
||||
request: Request,
|
||||
session: SessionDep,
|
||||
record: OtaRecord | None = None,
|
||||
) -> JavaJSONResponse:
|
||||
user = require_super_admin(request)
|
||||
if record is None:
|
||||
return error_response(request, 500, "固件信息不能为空")
|
||||
if is_blank(record.firmware_name):
|
||||
return error_response(request, 500, "固件名称不能为空")
|
||||
if is_blank(record.type):
|
||||
return error_response(request, 500, "固件类型不能为空")
|
||||
if is_blank(record.version):
|
||||
return error_response(request, 500, "版本号不能为空")
|
||||
try:
|
||||
await DeviceService(session).save_ota(record, user)
|
||||
return ok()
|
||||
except RuntimeError as exc:
|
||||
return error_response(request, 500, str(exc))
|
||||
|
||||
|
||||
@device_router.delete("/otaMag/{ota_id}")
|
||||
async def delete_ota(
|
||||
ota_id: str,
|
||||
request: Request,
|
||||
session: SessionDep,
|
||||
) -> JavaJSONResponse:
|
||||
require_super_admin(request)
|
||||
ids = ota_id.split(",") if ota_id else []
|
||||
if not ids:
|
||||
return error_response(request, 500, "删除的固件ID不能为空")
|
||||
await DeviceService(session).delete_ota(ids)
|
||||
return ok()
|
||||
|
||||
|
||||
@device_router.put("/otaMag/{ota_id}")
|
||||
async def update_ota(
|
||||
ota_id: str,
|
||||
request: Request,
|
||||
session: SessionDep,
|
||||
record: OtaRecord | None = None,
|
||||
) -> JavaJSONResponse:
|
||||
user = require_super_admin(request)
|
||||
if record is None:
|
||||
return error_response(request, 500, "固件信息不能为空")
|
||||
try:
|
||||
await DeviceService(session).update_ota(ota_id, record, user)
|
||||
return ok()
|
||||
except RuntimeError as exc:
|
||||
return error_response(request, 500, str(exc))
|
||||
|
||||
|
||||
def _validate_device_update(body: DeviceUpdateRequest, accept_language: str | None) -> str | None:
|
||||
language = resolve_language(accept_language)
|
||||
if body.auto_update is not None and body.auto_update < 0:
|
||||
return {
|
||||
"zh-CN": "最小不能小于0",
|
||||
"zh-TW": "必須大於或等於 0",
|
||||
"de-DE": "muss größer-gleich 0 sein",
|
||||
"pt-BR": "deve ser maior que ou igual à 0",
|
||||
}.get(language, "must be greater than or equal to 0")
|
||||
if body.auto_update is not None and body.auto_update > 1:
|
||||
return {
|
||||
"zh-CN": "最大不能超过1",
|
||||
"zh-TW": "必須小於或等於 1",
|
||||
"de-DE": "muss kleiner-gleich 1 sein",
|
||||
"pt-BR": "deve ser menor que ou igual à 1",
|
||||
}.get(language, "must be less than or equal to 1")
|
||||
if body.alias is not None and len(body.alias.encode("utf-16-le")) // 2 > 64:
|
||||
return {
|
||||
"zh-CN": "个数必须在0和64之间",
|
||||
"zh-TW": "大小必須在 0 和 64 之間",
|
||||
"de-DE": "Größe muss zwischen 0 und 64 sein",
|
||||
"pt-BR": "tamanho deve ser entre 0 e 64",
|
||||
}.get(language, "size must be between 0 and 64")
|
||||
return None
|
||||
@@ -0,0 +1,267 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
from typing import Annotated, Any
|
||||
|
||||
from fastapi import APIRouter, Depends, File, Form, Query, Request, UploadFile
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
|
||||
from app.core.database import get_db
|
||||
from app.core.errors import AppError
|
||||
from app.core.responses import JavaJSONResponse, envelope, ok
|
||||
from app.core.security import require_normal
|
||||
from app.repositories.knowledge import KnowledgeRepository
|
||||
from app.schemas.knowledge import DocumentBatchBody, KnowledgeBaseBody, RetrievalBody
|
||||
from app.services.knowledge import KnowledgeBaseService, KnowledgeDocumentService, dataset_dto
|
||||
|
||||
knowledge_router = APIRouter()
|
||||
|
||||
|
||||
def _base(session: AsyncSession) -> KnowledgeBaseService:
|
||||
return KnowledgeBaseService(KnowledgeRepository(session))
|
||||
|
||||
|
||||
def _documents(session: AsyncSession) -> KnowledgeDocumentService:
|
||||
return KnowledgeDocumentService(KnowledgeRepository(session))
|
||||
|
||||
|
||||
@knowledge_router.get("/datasets/rag-models")
|
||||
async def rag_models(request: Request, session: AsyncSession = Depends(get_db)) -> JavaJSONResponse:
|
||||
require_normal(request)
|
||||
return ok(await _base(session).rag_models())
|
||||
|
||||
|
||||
@knowledge_router.delete("/datasets/batch")
|
||||
async def delete_datasets_batch(
|
||||
request: Request, ids: str = Query(), session: AsyncSession = Depends(get_db)
|
||||
) -> JavaJSONResponse:
|
||||
user = require_normal(request)
|
||||
if not ids.strip():
|
||||
raise AppError(10003)
|
||||
await _base(session).batch_delete(
|
||||
ids.split(","), user, request.headers.get("Accept-Language")
|
||||
)
|
||||
return ok()
|
||||
|
||||
|
||||
@knowledge_router.get("/datasets")
|
||||
async def datasets_page(
|
||||
request: Request,
|
||||
name: str | None = None,
|
||||
page: int = 1,
|
||||
page_size: int = 10,
|
||||
session: AsyncSession = Depends(get_db),
|
||||
) -> JavaJSONResponse:
|
||||
return ok(
|
||||
await _base(session).page(
|
||||
require_normal(request),
|
||||
name,
|
||||
page,
|
||||
page_size,
|
||||
request.headers.get("Accept-Language"),
|
||||
)
|
||||
)
|
||||
|
||||
|
||||
@knowledge_router.post("/datasets")
|
||||
async def create_dataset(
|
||||
body: KnowledgeBaseBody, request: Request, session: AsyncSession = Depends(get_db)
|
||||
) -> JavaJSONResponse:
|
||||
return ok(await _base(session).create(body, require_normal(request)))
|
||||
|
||||
|
||||
@knowledge_router.get("/datasets/{dataset_id}")
|
||||
async def get_dataset(
|
||||
dataset_id: str, request: Request, session: AsyncSession = Depends(get_db)
|
||||
) -> JavaJSONResponse:
|
||||
return ok(dataset_dto(await _base(session).get_owned(dataset_id, require_normal(request))))
|
||||
|
||||
|
||||
@knowledge_router.put("/datasets/{dataset_id}")
|
||||
async def update_dataset(
|
||||
dataset_id: str,
|
||||
body: KnowledgeBaseBody,
|
||||
request: Request,
|
||||
session: AsyncSession = Depends(get_db),
|
||||
) -> JavaJSONResponse:
|
||||
return ok(await _base(session).update(dataset_id, body, require_normal(request)))
|
||||
|
||||
|
||||
@knowledge_router.delete("/datasets/{dataset_id}")
|
||||
async def delete_dataset(
|
||||
dataset_id: str, request: Request, session: AsyncSession = Depends(get_db)
|
||||
) -> JavaJSONResponse:
|
||||
await _base(session).delete(
|
||||
dataset_id, require_normal(request), request.headers.get("Accept-Language")
|
||||
)
|
||||
return ok()
|
||||
|
||||
|
||||
@knowledge_router.get("/datasets/{dataset_id}/documents/status/{status}")
|
||||
async def documents_by_status(
|
||||
dataset_id: str,
|
||||
status: str,
|
||||
request: Request,
|
||||
page: int = 1,
|
||||
page_size: int = 10,
|
||||
session: AsyncSession = Depends(get_db),
|
||||
) -> JavaJSONResponse:
|
||||
return ok(
|
||||
await _documents(session).page(
|
||||
dataset_id,
|
||||
require_normal(request),
|
||||
name=None,
|
||||
status=status,
|
||||
page=page,
|
||||
page_size=page_size,
|
||||
)
|
||||
)
|
||||
|
||||
|
||||
@knowledge_router.get("/datasets/{dataset_id}/documents")
|
||||
async def documents_page(
|
||||
dataset_id: str,
|
||||
request: Request,
|
||||
name: str | None = None,
|
||||
status: str | None = None,
|
||||
page: int = 1,
|
||||
page_size: int = 10,
|
||||
session: AsyncSession = Depends(get_db),
|
||||
) -> JavaJSONResponse:
|
||||
return ok(
|
||||
await _documents(session).page(
|
||||
dataset_id,
|
||||
require_normal(request),
|
||||
name=name,
|
||||
status=status,
|
||||
page=page,
|
||||
page_size=page_size,
|
||||
)
|
||||
)
|
||||
|
||||
|
||||
@knowledge_router.post("/datasets/{dataset_id}/documents")
|
||||
async def upload_document(
|
||||
dataset_id: str,
|
||||
request: Request,
|
||||
file: Annotated[UploadFile, File()],
|
||||
name: Annotated[str | None, Form()] = None,
|
||||
chunk_method: Annotated[str | None, Form(alias="chunkMethod")] = None,
|
||||
meta_fields: Annotated[str | None, Form(alias="metaFields")] = None,
|
||||
parser_config: Annotated[str | None, Form(alias="parserConfig")] = None,
|
||||
session: AsyncSession = Depends(get_db),
|
||||
) -> JavaJSONResponse:
|
||||
return ok(
|
||||
await _documents(session).upload(
|
||||
dataset_id,
|
||||
require_normal(request),
|
||||
file,
|
||||
name=name,
|
||||
meta_fields=_parse_form_json(meta_fields),
|
||||
chunk_method=chunk_method,
|
||||
parser_config=_parse_form_json(parser_config),
|
||||
)
|
||||
)
|
||||
|
||||
|
||||
@knowledge_router.delete("/datasets/{dataset_id}/documents")
|
||||
async def delete_documents(
|
||||
dataset_id: str,
|
||||
body: DocumentBatchBody,
|
||||
request: Request,
|
||||
session: AsyncSession = Depends(get_db),
|
||||
) -> JavaJSONResponse:
|
||||
await _documents(session).delete(
|
||||
dataset_id,
|
||||
body.ids,
|
||||
require_normal(request),
|
||||
request.headers.get("Accept-Language"),
|
||||
)
|
||||
return ok()
|
||||
|
||||
|
||||
@knowledge_router.delete("/datasets/{dataset_id}/documents/{document_id}")
|
||||
async def delete_document(
|
||||
dataset_id: str,
|
||||
document_id: str,
|
||||
request: Request,
|
||||
session: AsyncSession = Depends(get_db),
|
||||
) -> JavaJSONResponse:
|
||||
await _documents(session).delete(
|
||||
dataset_id,
|
||||
[document_id],
|
||||
require_normal(request),
|
||||
request.headers.get("Accept-Language"),
|
||||
)
|
||||
return ok()
|
||||
|
||||
|
||||
@knowledge_router.post("/datasets/{dataset_id}/chunks")
|
||||
async def parse_documents(
|
||||
dataset_id: str,
|
||||
body: dict[str, Any],
|
||||
request: Request,
|
||||
session: AsyncSession = Depends(get_db),
|
||||
) -> JavaJSONResponse:
|
||||
user = require_normal(request)
|
||||
# Java validates dataset existence/ownership before it reads document_ids.
|
||||
# A missing dataset must therefore win over the controller's empty-body
|
||||
# business error.
|
||||
await _base(session).get_owned(dataset_id, user)
|
||||
document_ids = body.get("document_ids")
|
||||
if document_ids is not None and not isinstance(document_ids, list):
|
||||
# Spring fails Map<String,List<String>> deserialization before entering
|
||||
# the controller, which is handled as the generic code-500 envelope.
|
||||
raise RuntimeError("document_ids must be an array")
|
||||
if not document_ids:
|
||||
return JavaJSONResponse(envelope(None, code=500, msg="document_ids参数不能为空"))
|
||||
success = await _documents(session).parse(dataset_id, document_ids, user)
|
||||
return ok() if success else JavaJSONResponse(
|
||||
envelope(None, code=500, msg="文档解析失败,文档可能正在处理中")
|
||||
)
|
||||
|
||||
|
||||
@knowledge_router.get("/datasets/{dataset_id}/documents/{document_id}/chunks")
|
||||
async def list_chunks(
|
||||
dataset_id: str,
|
||||
document_id: str,
|
||||
request: Request,
|
||||
page: int = 1,
|
||||
page_size: int = 10,
|
||||
keywords: str | None = None,
|
||||
id: str | None = None, # noqa: A002 - exact Java query parameter
|
||||
session: AsyncSession = Depends(get_db),
|
||||
) -> JavaJSONResponse:
|
||||
return ok(
|
||||
await _documents(session).chunks(
|
||||
dataset_id,
|
||||
document_id,
|
||||
require_normal(request),
|
||||
page=page,
|
||||
page_size=page_size,
|
||||
keywords=keywords,
|
||||
chunk_id=id,
|
||||
)
|
||||
)
|
||||
|
||||
|
||||
@knowledge_router.post("/datasets/{dataset_id}/retrieval-test")
|
||||
async def retrieval_test(
|
||||
dataset_id: str,
|
||||
body: RetrievalBody,
|
||||
request: Request,
|
||||
session: AsyncSession = Depends(get_db),
|
||||
) -> JavaJSONResponse:
|
||||
return ok(await _documents(session).retrieval(dataset_id, body, require_normal(request)))
|
||||
|
||||
|
||||
def _parse_form_json(value: str | None) -> dict[str, Any] | None:
|
||||
if value is None:
|
||||
return None
|
||||
try:
|
||||
result = json.loads(value)
|
||||
except json.JSONDecodeError as exc:
|
||||
raise RuntimeError(f"解析JSON字符串失败: {value}") from exc
|
||||
if not isinstance(result, dict):
|
||||
raise RuntimeError(f"解析JSON字符串失败: {value}")
|
||||
return dict(result)
|
||||
@@ -0,0 +1,178 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from fastapi import APIRouter, Depends, Query, Request
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
|
||||
from app.core.database import get_db
|
||||
from app.core.responses import JavaJSONResponse, envelope, ok
|
||||
from app.core.security import require_normal, require_super_admin
|
||||
from app.repositories.config import ConfigRepository
|
||||
from app.repositories.model import ModelRepository
|
||||
from app.schemas.model import ModelConfigBody, ModelProviderBody
|
||||
from app.services.config import ConfigService
|
||||
from app.services.model import ModelProviderService, ModelService
|
||||
|
||||
model_router = APIRouter()
|
||||
|
||||
|
||||
def _models(session: AsyncSession) -> ModelService:
|
||||
return ModelService(ModelRepository(session))
|
||||
|
||||
|
||||
def _providers(session: AsyncSession) -> ModelProviderService:
|
||||
return ModelProviderService(ModelRepository(session))
|
||||
|
||||
|
||||
async def _refresh_server_config(session: AsyncSession) -> None:
|
||||
await ConfigService(ConfigRepository(session)).get_config(use_cache=False)
|
||||
|
||||
|
||||
@model_router.get("/models/names")
|
||||
async def model_names(
|
||||
request: Request,
|
||||
model_type: str = Query(alias="modelType"),
|
||||
model_name: str | None = Query(default=None, alias="modelName"),
|
||||
session: AsyncSession = Depends(get_db),
|
||||
) -> JavaJSONResponse:
|
||||
require_normal(request)
|
||||
return ok(await _models(session).names(model_type, model_name))
|
||||
|
||||
|
||||
@model_router.get("/models/llm/names")
|
||||
async def llm_names(
|
||||
request: Request,
|
||||
model_name: str | None = Query(default=None, alias="modelName"),
|
||||
session: AsyncSession = Depends(get_db),
|
||||
) -> JavaJSONResponse:
|
||||
require_normal(request)
|
||||
return ok(await _models(session).llm_names(model_name))
|
||||
|
||||
|
||||
@model_router.get("/models/list")
|
||||
async def model_list(
|
||||
request: Request,
|
||||
model_type: str = Query(alias="modelType"),
|
||||
model_name: str | None = Query(default=None, alias="modelName"),
|
||||
page: str = "0",
|
||||
limit: str = "10",
|
||||
session: AsyncSession = Depends(get_db),
|
||||
) -> JavaJSONResponse:
|
||||
require_super_admin(request)
|
||||
return ok(await _models(session).model_page(model_type, model_name, page, limit))
|
||||
|
||||
|
||||
@model_router.get("/models/provider/plugin/names")
|
||||
async def plugin_names(request: Request, session: AsyncSession = Depends(get_db)) -> JavaJSONResponse:
|
||||
user = require_normal(request)
|
||||
return ok(await ModelRepository(session).list_plugins_for_user(user.id))
|
||||
|
||||
|
||||
@model_router.get("/models/provider")
|
||||
async def provider_list(
|
||||
request: Request,
|
||||
model_type: str | None = Query(default=None, alias="modelType"),
|
||||
name: str | None = None,
|
||||
page: str = "0",
|
||||
limit: str = "10",
|
||||
session: AsyncSession = Depends(get_db),
|
||||
) -> JavaJSONResponse:
|
||||
require_super_admin(request)
|
||||
return ok(await _providers(session).page(model_type, name, page, limit))
|
||||
|
||||
|
||||
@model_router.post("/models/provider")
|
||||
async def provider_add(
|
||||
body: ModelProviderBody, request: Request, session: AsyncSession = Depends(get_db)
|
||||
) -> JavaJSONResponse:
|
||||
return ok(await _providers(session).add(body, require_super_admin(request)))
|
||||
|
||||
|
||||
@model_router.put("/models/provider")
|
||||
async def provider_edit(
|
||||
body: ModelProviderBody, request: Request, session: AsyncSession = Depends(get_db)
|
||||
) -> JavaJSONResponse:
|
||||
return ok(await _providers(session).edit(body, require_super_admin(request)))
|
||||
|
||||
|
||||
@model_router.post("/models/provider/delete")
|
||||
async def provider_delete(
|
||||
ids: list[str], request: Request, session: AsyncSession = Depends(get_db)
|
||||
) -> JavaJSONResponse:
|
||||
require_super_admin(request)
|
||||
await _providers(session).delete(ids)
|
||||
return ok()
|
||||
|
||||
|
||||
@model_router.get("/models/{model_type}/provideTypes")
|
||||
async def provider_types(
|
||||
model_type: str, request: Request, session: AsyncSession = Depends(get_db)
|
||||
) -> JavaJSONResponse:
|
||||
require_super_admin(request)
|
||||
return ok(await ModelRepository(session).list_providers_by_type(model_type))
|
||||
|
||||
|
||||
@model_router.post("/models/{model_type}/{provide_code}")
|
||||
async def model_add(
|
||||
model_type: str,
|
||||
provide_code: str,
|
||||
body: ModelConfigBody,
|
||||
request: Request,
|
||||
session: AsyncSession = Depends(get_db),
|
||||
) -> JavaJSONResponse:
|
||||
require_super_admin(request)
|
||||
result = await _models(session).add(model_type, provide_code, body)
|
||||
await _refresh_server_config(session)
|
||||
return ok(result)
|
||||
|
||||
|
||||
@model_router.put("/models/enable/{model_id}/{status}")
|
||||
async def model_enable(
|
||||
model_id: str, status: int, request: Request, session: AsyncSession = Depends(get_db)
|
||||
) -> JavaJSONResponse:
|
||||
require_super_admin(request)
|
||||
message = await _models(session).enable(model_id, status)
|
||||
return JavaJSONResponse(envelope(None, code=500, msg=message)) if message else ok()
|
||||
|
||||
|
||||
@model_router.put("/models/{model_type}/{provide_code}/{model_id}")
|
||||
async def model_edit(
|
||||
model_type: str,
|
||||
provide_code: str,
|
||||
model_id: str,
|
||||
body: ModelConfigBody,
|
||||
request: Request,
|
||||
session: AsyncSession = Depends(get_db),
|
||||
) -> JavaJSONResponse:
|
||||
require_super_admin(request)
|
||||
result = await _models(session).edit(model_type, provide_code, model_id, body)
|
||||
await _refresh_server_config(session)
|
||||
return ok(result)
|
||||
|
||||
|
||||
@model_router.put("/models/default/{model_id}")
|
||||
async def model_default(
|
||||
model_id: str, request: Request, session: AsyncSession = Depends(get_db)
|
||||
) -> JavaJSONResponse:
|
||||
require_super_admin(request)
|
||||
message = await _models(session).set_default(model_id)
|
||||
if message:
|
||||
return JavaJSONResponse(envelope(None, code=500, msg=message))
|
||||
await _refresh_server_config(session)
|
||||
return ok()
|
||||
|
||||
|
||||
@model_router.get("/models/{model_id}")
|
||||
async def model_get(
|
||||
model_id: str, request: Request, session: AsyncSession = Depends(get_db)
|
||||
) -> JavaJSONResponse:
|
||||
require_super_admin(request)
|
||||
return ok(await _models(session).get_model(model_id))
|
||||
|
||||
|
||||
@model_router.delete("/models/{model_id}")
|
||||
async def model_delete(
|
||||
model_id: str, request: Request, session: AsyncSession = Depends(get_db)
|
||||
) -> JavaJSONResponse:
|
||||
require_super_admin(request)
|
||||
await _models(session).delete(model_id)
|
||||
return ok()
|
||||
@@ -0,0 +1,111 @@
|
||||
# ruff: noqa: B008
|
||||
# FastAPI evaluates dependency marker defaults intentionally when registering routes.
|
||||
from __future__ import annotations
|
||||
|
||||
from fastapi import APIRouter, Depends, Query, Request
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
from starlette.responses import Response
|
||||
|
||||
from app.core.database import get_db
|
||||
from app.core.errors import AppError
|
||||
from app.core.responses import JavaJSONResponse, ok
|
||||
from app.core.security import require_normal
|
||||
from app.repositories.security import SecurityRepository
|
||||
from app.schemas.security import (
|
||||
LoginRequest,
|
||||
PasswordChangeRequest,
|
||||
RetrievePasswordRequest,
|
||||
SmsVerificationRequest,
|
||||
)
|
||||
from app.services.security import CaptchaService, SecurityService
|
||||
|
||||
security_router = APIRouter()
|
||||
|
||||
|
||||
def _service(session: AsyncSession) -> SecurityService:
|
||||
return SecurityService(SecurityRepository(session))
|
||||
|
||||
|
||||
@security_router.get("/user/captcha")
|
||||
async def captcha(uuid: str | None = Query(default=None)) -> Response:
|
||||
if uuid is None or not uuid.strip():
|
||||
raise AppError(10006)
|
||||
content = await CaptchaService().create(uuid)
|
||||
return Response(
|
||||
content,
|
||||
media_type="image/gif",
|
||||
headers={
|
||||
"Pragma": "No-cache",
|
||||
"Cache-Control": "no-cache",
|
||||
"Expires": "Thu, 01 Jan 1970 00:00:00 GMT",
|
||||
},
|
||||
)
|
||||
|
||||
|
||||
@security_router.post("/user/smsVerification")
|
||||
async def sms_verification(dto: SmsVerificationRequest, session: AsyncSession = Depends(get_db)) -> JavaJSONResponse:
|
||||
await _service(session).send_sms_verification(dto)
|
||||
return ok()
|
||||
|
||||
|
||||
@security_router.post("/user/login")
|
||||
async def login(
|
||||
dto: LoginRequest,
|
||||
request: Request,
|
||||
session: AsyncSession = Depends(get_db),
|
||||
) -> JavaJSONResponse:
|
||||
return ok(await _service(session).login(dto, request))
|
||||
|
||||
|
||||
@security_router.post("/user/register")
|
||||
async def register(dto: LoginRequest, session: AsyncSession = Depends(get_db)) -> JavaJSONResponse:
|
||||
await _service(session).register(dto)
|
||||
return ok()
|
||||
|
||||
|
||||
@security_router.get("/user/info")
|
||||
async def info(request: Request) -> JavaJSONResponse:
|
||||
user = require_normal(request)
|
||||
return ok(
|
||||
{
|
||||
"id": user.id,
|
||||
"username": user.username,
|
||||
"superAdmin": user.super_admin,
|
||||
"token": user.token,
|
||||
"status": user.status,
|
||||
}
|
||||
)
|
||||
|
||||
|
||||
@security_router.put("/user/change-password")
|
||||
async def change_password(
|
||||
dto: PasswordChangeRequest,
|
||||
request: Request,
|
||||
session: AsyncSession = Depends(get_db),
|
||||
) -> JavaJSONResponse:
|
||||
await _service(session).change_password(
|
||||
require_normal(request),
|
||||
dto,
|
||||
request.headers.get("Accept-Language"),
|
||||
)
|
||||
return ok()
|
||||
|
||||
|
||||
@security_router.put("/user/retrieve-password")
|
||||
async def retrieve_password(
|
||||
dto: RetrievePasswordRequest,
|
||||
request: Request,
|
||||
session: AsyncSession = Depends(get_db),
|
||||
) -> JavaJSONResponse:
|
||||
await _service(session).retrieve_password(dto, request.headers.get("Accept-Language"))
|
||||
return ok()
|
||||
|
||||
|
||||
@security_router.get("/user/pub-config")
|
||||
async def public_config(session: AsyncSession = Depends(get_db)) -> JavaJSONResponse:
|
||||
return ok(await _service(session).public_config())
|
||||
|
||||
|
||||
@security_router.get("/api/ping")
|
||||
async def api_ping() -> JavaJSONResponse:
|
||||
return ok("pong")
|
||||
@@ -0,0 +1,334 @@
|
||||
# ruff: noqa: B008
|
||||
# FastAPI evaluates dependency and body marker defaults intentionally when registering routes.
|
||||
from __future__ import annotations
|
||||
|
||||
from fastapi import APIRouter, Body, Depends, Query, Request
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
|
||||
from app.core.database import get_db
|
||||
from app.core.errors import AppError
|
||||
from app.core.responses import JavaJSONResponse, envelope, ok
|
||||
from app.core.security import require_normal, require_super_admin
|
||||
from app.repositories.sys import SysRepository
|
||||
from app.schemas.sys import DictDataPayload, DictTypePayload, EmitServerActionRequest, SysParamPayload
|
||||
from app.services.sys import AdminService, DictService, ServerActionService, SysParamService
|
||||
|
||||
sys_router = APIRouter()
|
||||
|
||||
|
||||
def _repository(session: AsyncSession) -> SysRepository:
|
||||
return SysRepository(session)
|
||||
|
||||
|
||||
def _admin(session: AsyncSession) -> AdminService:
|
||||
return AdminService(_repository(session))
|
||||
|
||||
|
||||
def _params(session: AsyncSession) -> SysParamService:
|
||||
return SysParamService(_repository(session))
|
||||
|
||||
|
||||
def _dict(session: AsyncSession) -> DictService:
|
||||
return DictService(_repository(session))
|
||||
|
||||
|
||||
async def _refresh_server_config(session: AsyncSession) -> None:
|
||||
from app.repositories.config import ConfigRepository
|
||||
from app.services.config import ConfigService
|
||||
|
||||
await ConfigService(ConfigRepository(session)).get_config(use_cache=False)
|
||||
|
||||
|
||||
@sys_router.get("/admin/users")
|
||||
async def page_users(
|
||||
request: Request,
|
||||
mobile: str | None = None,
|
||||
page: str = Query(default="1"),
|
||||
limit: str = Query(default="10"),
|
||||
session: AsyncSession = Depends(get_db),
|
||||
) -> JavaJSONResponse:
|
||||
require_super_admin(request)
|
||||
try:
|
||||
current, size = int(page), int(limit)
|
||||
except ValueError as exc:
|
||||
# Java parses these Map-backed values inside the service; malformed
|
||||
# numbers therefore reach its generic code=500 handler rather than
|
||||
# Bean Validation.
|
||||
raise AppError(500, "排序值不能小于0") from exc
|
||||
return ok(await _admin(session).page_users(mobile=mobile, page=current, limit=size))
|
||||
|
||||
|
||||
@sys_router.put("/admin/users/{user_id}")
|
||||
async def reset_user_password(
|
||||
user_id: int,
|
||||
request: Request,
|
||||
session: AsyncSession = Depends(get_db),
|
||||
) -> JavaJSONResponse:
|
||||
user = require_super_admin(request)
|
||||
return ok(await _admin(session).reset_password(user_id, user))
|
||||
|
||||
|
||||
@sys_router.delete("/admin/users/{user_id}")
|
||||
async def delete_user(
|
||||
user_id: int,
|
||||
request: Request,
|
||||
session: AsyncSession = Depends(get_db),
|
||||
) -> JavaJSONResponse:
|
||||
require_super_admin(request)
|
||||
await _admin(session).delete_user(user_id)
|
||||
return ok()
|
||||
|
||||
|
||||
@sys_router.put("/admin/users/changeStatus/{status}")
|
||||
async def change_user_status(
|
||||
status: int,
|
||||
request: Request,
|
||||
user_ids: list[str] = Body(),
|
||||
session: AsyncSession = Depends(get_db),
|
||||
) -> JavaJSONResponse:
|
||||
user = require_super_admin(request)
|
||||
await _admin(session).change_status(status, user_ids, user)
|
||||
return ok()
|
||||
|
||||
|
||||
@sys_router.get("/admin/device/all")
|
||||
async def page_all_devices(
|
||||
request: Request,
|
||||
keywords: str | None = None,
|
||||
page: int = Query(default=1, ge=0),
|
||||
limit: int = Query(default=10, ge=0),
|
||||
session: AsyncSession = Depends(get_db),
|
||||
) -> JavaJSONResponse:
|
||||
require_super_admin(request)
|
||||
return ok(await _admin(session).page_devices(keywords=keywords, page=page, limit=limit))
|
||||
|
||||
|
||||
@sys_router.get("/admin/server/server-list")
|
||||
async def websocket_server_list(request: Request, session: AsyncSession = Depends(get_db)) -> JavaJSONResponse:
|
||||
require_super_admin(request)
|
||||
params = _params(session)
|
||||
return ok(await ServerActionService(params).server_list())
|
||||
|
||||
|
||||
@sys_router.post("/admin/server/emit-action")
|
||||
async def emit_server_action(
|
||||
dto: EmitServerActionRequest,
|
||||
request: Request,
|
||||
session: AsyncSession = Depends(get_db),
|
||||
) -> JavaJSONResponse:
|
||||
require_super_admin(request)
|
||||
return ok(await ServerActionService(_params(session)).emit(dto))
|
||||
|
||||
|
||||
@sys_router.get("/admin/params/page")
|
||||
async def page_params(
|
||||
request: Request,
|
||||
page: int = Query(default=1, ge=0),
|
||||
limit: int = Query(default=10, ge=0),
|
||||
order_field: str | None = Query(default=None, alias="orderField"),
|
||||
order: str | None = None,
|
||||
param_code: str | None = Query(default=None, alias="paramCode"),
|
||||
session: AsyncSession = Depends(get_db),
|
||||
) -> JavaJSONResponse:
|
||||
require_super_admin(request)
|
||||
return ok(
|
||||
await _params(session).page(
|
||||
param_code=param_code,
|
||||
page=page,
|
||||
limit=limit,
|
||||
order_field=order_field,
|
||||
order=order,
|
||||
)
|
||||
)
|
||||
|
||||
|
||||
@sys_router.get("/admin/params/{param_id}")
|
||||
async def get_param(
|
||||
param_id: int,
|
||||
request: Request,
|
||||
session: AsyncSession = Depends(get_db),
|
||||
) -> JavaJSONResponse:
|
||||
require_super_admin(request)
|
||||
return ok(await _params(session).get(param_id))
|
||||
|
||||
|
||||
@sys_router.post("/admin/params")
|
||||
async def save_param(
|
||||
dto: SysParamPayload,
|
||||
request: Request,
|
||||
session: AsyncSession = Depends(get_db),
|
||||
) -> JavaJSONResponse:
|
||||
await _params(session).save(
|
||||
dto,
|
||||
require_super_admin(request),
|
||||
request.headers.get("Accept-Language"),
|
||||
)
|
||||
await _refresh_server_config(session)
|
||||
return ok()
|
||||
|
||||
|
||||
@sys_router.put("/admin/params")
|
||||
async def update_param(
|
||||
dto: SysParamPayload,
|
||||
request: Request,
|
||||
session: AsyncSession = Depends(get_db),
|
||||
) -> JavaJSONResponse:
|
||||
await _params(session).update(
|
||||
dto,
|
||||
require_super_admin(request),
|
||||
request.headers.get("Accept-Language"),
|
||||
)
|
||||
await _refresh_server_config(session)
|
||||
return ok()
|
||||
|
||||
|
||||
@sys_router.post("/admin/params/delete")
|
||||
async def delete_params(
|
||||
request: Request,
|
||||
ids: list[str] = Body(),
|
||||
session: AsyncSession = Depends(get_db),
|
||||
) -> JavaJSONResponse:
|
||||
require_super_admin(request)
|
||||
await _params(session).delete(ids)
|
||||
await _refresh_server_config(session)
|
||||
return ok()
|
||||
|
||||
|
||||
@sys_router.get("/admin/dict/type/page")
|
||||
async def page_dict_types(
|
||||
request: Request,
|
||||
dict_type: str | None = Query(default=None, alias="dictType"),
|
||||
dict_name: str | None = Query(default=None, alias="dictName"),
|
||||
page: int = Query(default=1, ge=0),
|
||||
limit: int = Query(default=10, ge=0),
|
||||
session: AsyncSession = Depends(get_db),
|
||||
) -> JavaJSONResponse:
|
||||
require_super_admin(request)
|
||||
return ok(
|
||||
await _dict(session).page_types(
|
||||
dict_type=dict_type,
|
||||
dict_name=dict_name,
|
||||
page=page,
|
||||
limit=limit,
|
||||
)
|
||||
)
|
||||
|
||||
|
||||
@sys_router.get("/admin/dict/type/{type_id}")
|
||||
async def get_dict_type(
|
||||
type_id: int,
|
||||
request: Request,
|
||||
session: AsyncSession = Depends(get_db),
|
||||
) -> JavaJSONResponse:
|
||||
require_super_admin(request)
|
||||
return ok(await _dict(session).get_type(type_id))
|
||||
|
||||
|
||||
@sys_router.post("/admin/dict/type/save")
|
||||
async def save_dict_type(
|
||||
dto: DictTypePayload,
|
||||
request: Request,
|
||||
session: AsyncSession = Depends(get_db),
|
||||
) -> JavaJSONResponse:
|
||||
await _dict(session).save_type(dto, require_super_admin(request))
|
||||
return ok()
|
||||
|
||||
|
||||
@sys_router.put("/admin/dict/type/update")
|
||||
async def update_dict_type(
|
||||
dto: DictTypePayload,
|
||||
request: Request,
|
||||
session: AsyncSession = Depends(get_db),
|
||||
) -> JavaJSONResponse:
|
||||
await _dict(session).update_type(dto, require_super_admin(request))
|
||||
return ok()
|
||||
|
||||
|
||||
@sys_router.post("/admin/dict/type/delete")
|
||||
async def delete_dict_types(
|
||||
request: Request,
|
||||
ids: list[int] = Body(),
|
||||
session: AsyncSession = Depends(get_db),
|
||||
) -> JavaJSONResponse:
|
||||
require_super_admin(request)
|
||||
await _dict(session).delete_types(ids)
|
||||
return ok()
|
||||
|
||||
|
||||
@sys_router.get("/admin/dict/data/page")
|
||||
async def page_dict_data(
|
||||
request: Request,
|
||||
dict_type_id: str | None = Query(default=None, alias="dictTypeId"),
|
||||
dict_label: str | None = Query(default=None, alias="dictLabel"),
|
||||
dict_value: str | None = Query(default=None, alias="dictValue"),
|
||||
page: int = Query(default=1, ge=0),
|
||||
limit: int = Query(default=10, ge=0),
|
||||
session: AsyncSession = Depends(get_db),
|
||||
) -> JavaJSONResponse:
|
||||
require_super_admin(request)
|
||||
if dict_type_id is None or not dict_type_id:
|
||||
return JavaJSONResponse(envelope(None, code=500, msg="dictTypeId不能为空"))
|
||||
try:
|
||||
parsed_type_id = int(dict_type_id)
|
||||
except ValueError as exc:
|
||||
raise AppError(500) from exc
|
||||
return ok(
|
||||
await _dict(session).page_data(
|
||||
dict_type_id=parsed_type_id,
|
||||
dict_label=dict_label,
|
||||
dict_value=dict_value,
|
||||
page=page,
|
||||
limit=limit,
|
||||
)
|
||||
)
|
||||
|
||||
|
||||
@sys_router.get("/admin/dict/data/type/{dict_type}")
|
||||
async def dict_items(
|
||||
dict_type: str,
|
||||
request: Request,
|
||||
session: AsyncSession = Depends(get_db),
|
||||
) -> JavaJSONResponse:
|
||||
require_normal(request)
|
||||
return ok(await _dict(session).items(dict_type))
|
||||
|
||||
|
||||
@sys_router.get("/admin/dict/data/{data_id}")
|
||||
async def get_dict_data(
|
||||
data_id: int,
|
||||
request: Request,
|
||||
session: AsyncSession = Depends(get_db),
|
||||
) -> JavaJSONResponse:
|
||||
require_super_admin(request)
|
||||
return ok(await _dict(session).get_data(data_id))
|
||||
|
||||
|
||||
@sys_router.post("/admin/dict/data/save")
|
||||
async def save_dict_data(
|
||||
dto: DictDataPayload,
|
||||
request: Request,
|
||||
session: AsyncSession = Depends(get_db),
|
||||
) -> JavaJSONResponse:
|
||||
await _dict(session).save_data(dto, require_super_admin(request))
|
||||
return ok()
|
||||
|
||||
|
||||
@sys_router.put("/admin/dict/data/update")
|
||||
async def update_dict_data(
|
||||
dto: DictDataPayload,
|
||||
request: Request,
|
||||
session: AsyncSession = Depends(get_db),
|
||||
) -> JavaJSONResponse:
|
||||
await _dict(session).update_data(dto, require_super_admin(request))
|
||||
return ok()
|
||||
|
||||
|
||||
@sys_router.post("/admin/dict/data/delete")
|
||||
async def delete_dict_data(
|
||||
request: Request,
|
||||
ids: list[int] = Body(),
|
||||
session: AsyncSession = Depends(get_db),
|
||||
) -> JavaJSONResponse:
|
||||
require_super_admin(request)
|
||||
await _dict(session).delete_data(ids)
|
||||
return ok()
|
||||
@@ -0,0 +1,74 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from fastapi import APIRouter, Depends, Query, Request
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
|
||||
from app.core.database import get_db
|
||||
from app.core.responses import JavaJSONResponse, ok
|
||||
from app.core.security import require_normal, require_super_admin
|
||||
from app.repositories.timbre import TimbreRepository
|
||||
from app.schemas.timbre import TimbreBody
|
||||
from app.services.timbre import TimbreService
|
||||
|
||||
timbre_router = APIRouter()
|
||||
|
||||
|
||||
def _service(session: AsyncSession) -> TimbreService:
|
||||
return TimbreService(TimbreRepository(session))
|
||||
|
||||
|
||||
@timbre_router.get("/ttsVoice")
|
||||
async def timbre_page(
|
||||
request: Request,
|
||||
tts_model_id: str | None = Query(default=None, alias="ttsModelId"),
|
||||
name: str | None = None,
|
||||
page: str | None = None,
|
||||
limit: str | None = None,
|
||||
session: AsyncSession = Depends(get_db),
|
||||
) -> JavaJSONResponse:
|
||||
require_super_admin(request)
|
||||
return ok(
|
||||
await _service(session).page(
|
||||
tts_model_id, name, page, limit, request.headers.get("Accept-Language")
|
||||
)
|
||||
)
|
||||
|
||||
|
||||
@timbre_router.post("/ttsVoice")
|
||||
async def timbre_save(
|
||||
body: TimbreBody, request: Request, session: AsyncSession = Depends(get_db)
|
||||
) -> JavaJSONResponse:
|
||||
await _service(session).save(
|
||||
body, require_super_admin(request), request.headers.get("Accept-Language")
|
||||
)
|
||||
return ok()
|
||||
|
||||
|
||||
@timbre_router.put("/ttsVoice/{timbre_id}")
|
||||
async def timbre_update(
|
||||
timbre_id: str, body: TimbreBody, request: Request, session: AsyncSession = Depends(get_db)
|
||||
) -> JavaJSONResponse:
|
||||
await _service(session).update(
|
||||
timbre_id, body, require_super_admin(request), request.headers.get("Accept-Language")
|
||||
)
|
||||
return ok()
|
||||
|
||||
|
||||
@timbre_router.post("/ttsVoice/delete")
|
||||
async def timbre_delete(
|
||||
ids: list[str], request: Request, session: AsyncSession = Depends(get_db)
|
||||
) -> JavaJSONResponse:
|
||||
require_super_admin(request)
|
||||
await _service(session).delete(ids)
|
||||
return ok()
|
||||
|
||||
|
||||
@timbre_router.get("/models/{model_id}/voices")
|
||||
async def model_voices(
|
||||
model_id: str,
|
||||
request: Request,
|
||||
voice_name: str | None = Query(default=None, alias="voiceName"),
|
||||
session: AsyncSession = Depends(get_db),
|
||||
) -> JavaJSONResponse:
|
||||
user = require_normal(request)
|
||||
return ok(await _service(session).voices(model_id, voice_name, user, request.headers.get("Accept-Language")))
|
||||
@@ -0,0 +1,222 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Annotated, Any
|
||||
|
||||
from fastapi import APIRouter, Depends, File, Form, Request, UploadFile
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
from starlette.responses import Response
|
||||
|
||||
from app.core.database import get_db
|
||||
from app.core.errors import AppError
|
||||
from app.core.i18n import message_for
|
||||
from app.core.responses import JavaJSONResponse, error_response, ok
|
||||
from app.core.security import require_normal, require_super_admin
|
||||
from app.schemas.voiceclone import VoiceCloneRenameRequest, VoiceCloneTrainRequest, VoiceResourceCreateRequest
|
||||
from app.services.voiceclone import VoiceCloneService
|
||||
|
||||
voiceclone_router = APIRouter()
|
||||
SessionDep = Annotated[AsyncSession, Depends(get_db)]
|
||||
VoiceFile = Annotated[UploadFile, File(alias="voiceFile")]
|
||||
VoiceIdForm = Annotated[str, Form(alias="id")]
|
||||
|
||||
|
||||
def _query_map(request: Request) -> dict[str, Any]:
|
||||
result: dict[str, Any] = {}
|
||||
for key, value in request.query_params.multi_items():
|
||||
if key in result:
|
||||
previous = result[key]
|
||||
result[key] = [*previous, value] if isinstance(previous, list) else [previous, value]
|
||||
else:
|
||||
result[key] = value
|
||||
return result
|
||||
|
||||
|
||||
# Static voiceResource paths deliberately precede /voiceResource/{id}.
|
||||
@voiceclone_router.get("/voiceResource/ttsPlatforms")
|
||||
async def tts_platforms(
|
||||
request: Request,
|
||||
session: SessionDep,
|
||||
) -> JavaJSONResponse:
|
||||
require_super_admin(request)
|
||||
return ok(await VoiceCloneService(session).tts_platforms())
|
||||
|
||||
|
||||
@voiceclone_router.get("/voiceResource/user/{user_id}")
|
||||
async def voice_resources_by_user(
|
||||
user_id: int,
|
||||
request: Request,
|
||||
session: SessionDep,
|
||||
) -> JavaJSONResponse:
|
||||
require_normal(request)
|
||||
return ok(await VoiceCloneService(session).get_by_user(user_id))
|
||||
|
||||
|
||||
@voiceclone_router.get("/voiceResource")
|
||||
async def page_voice_resources(
|
||||
request: Request,
|
||||
session: SessionDep,
|
||||
) -> JavaJSONResponse:
|
||||
require_super_admin(request)
|
||||
return ok(await VoiceCloneService(session).page(_query_map(request)))
|
||||
|
||||
|
||||
@voiceclone_router.get("/voiceResource/{voice_id}")
|
||||
async def get_voice_resource(
|
||||
voice_id: str,
|
||||
request: Request,
|
||||
session: SessionDep,
|
||||
) -> JavaJSONResponse:
|
||||
require_super_admin(request)
|
||||
return ok(await VoiceCloneService(session).get_detail(voice_id))
|
||||
|
||||
|
||||
@voiceclone_router.post("/voiceResource")
|
||||
async def create_voice_resource(
|
||||
request: Request,
|
||||
session: SessionDep,
|
||||
body: VoiceResourceCreateRequest | None = None,
|
||||
) -> JavaJSONResponse:
|
||||
user = require_super_admin(request)
|
||||
if body is None:
|
||||
return error_response(request, 10145)
|
||||
if body.model_id is None or body.model_id == "":
|
||||
return error_response(request, 10146)
|
||||
if not body.voice_ids:
|
||||
return error_response(request, 10147)
|
||||
if body.user_id is None:
|
||||
return error_response(request, 10148)
|
||||
try:
|
||||
await VoiceCloneService(session).create_resources(body, actor=user)
|
||||
return ok()
|
||||
except AppError:
|
||||
raise
|
||||
except RuntimeError as exc:
|
||||
return error_response(request, 10065, str(exc))
|
||||
|
||||
|
||||
@voiceclone_router.delete("/voiceResource/{voice_id}")
|
||||
async def delete_voice_resource(
|
||||
voice_id: str,
|
||||
request: Request,
|
||||
session: SessionDep,
|
||||
) -> JavaJSONResponse:
|
||||
require_super_admin(request)
|
||||
ids = voice_id.split(",") if voice_id else []
|
||||
if not ids:
|
||||
return error_response(request, 10149)
|
||||
await VoiceCloneService(session).delete(ids)
|
||||
return ok()
|
||||
|
||||
|
||||
@voiceclone_router.get("/voiceClone")
|
||||
async def page_voice_clones(
|
||||
request: Request,
|
||||
session: SessionDep,
|
||||
) -> JavaJSONResponse:
|
||||
user = require_normal(request)
|
||||
return ok(await VoiceCloneService(session).page(_query_map(request), user_id=user.id))
|
||||
|
||||
|
||||
@voiceclone_router.post("/voiceClone/upload")
|
||||
async def upload_voice_clone(
|
||||
request: Request,
|
||||
session: SessionDep,
|
||||
voice_file: VoiceFile,
|
||||
voice_id: VoiceIdForm = "",
|
||||
) -> JavaJSONResponse:
|
||||
user = require_normal(request)
|
||||
service = VoiceCloneService(session)
|
||||
try:
|
||||
content = await voice_file.read()
|
||||
if not content:
|
||||
return error_response(request, 10140)
|
||||
content_type = voice_file.content_type
|
||||
if content_type is None or not content_type.startswith("audio/"):
|
||||
return error_response(request, 10141)
|
||||
filename = voice_file.filename
|
||||
if filename is None or "." not in filename:
|
||||
raise RuntimeError("文件名缺少扩展名")
|
||||
extension = filename[filename.rfind(".") :].lower()
|
||||
if extension not in {".mp3", ".wav"}:
|
||||
return error_response(request, 500, "只允许上传.mp3和.wav格式的文件")
|
||||
if len(content) > 10 * 1024 * 1024:
|
||||
return error_response(request, 10142)
|
||||
await service.check_permission(voice_id, user)
|
||||
await service.upload_voice(voice_id, content)
|
||||
return ok()
|
||||
except Exception as exc:
|
||||
if isinstance(exc, AppError):
|
||||
message = exc.message or message_for(exc.code, request.headers.get("Accept-Language"))
|
||||
else:
|
||||
message = str(exc)
|
||||
return error_response(request, 10143, message)
|
||||
|
||||
|
||||
@voiceclone_router.post("/voiceClone/updateName")
|
||||
async def update_voice_clone_name(
|
||||
body: VoiceCloneRenameRequest,
|
||||
request: Request,
|
||||
session: SessionDep,
|
||||
) -> JavaJSONResponse:
|
||||
user = require_normal(request)
|
||||
if body.id is None or body.id == "":
|
||||
return error_response(request, 10006)
|
||||
if body.name is None or body.name == "":
|
||||
return error_response(request, 10181)
|
||||
service = VoiceCloneService(session)
|
||||
try:
|
||||
await service.check_permission(body.id, user)
|
||||
await service.rename(body.id or "", body.name or "")
|
||||
return ok()
|
||||
except Exception as exc:
|
||||
if isinstance(exc, AppError):
|
||||
message = exc.message or message_for(exc.code, request.headers.get("Accept-Language"))
|
||||
else:
|
||||
message = str(exc)
|
||||
return error_response(request, 10066, message)
|
||||
|
||||
|
||||
@voiceclone_router.post("/voiceClone/audio/{voice_id}")
|
||||
async def get_voice_clone_audio_id(
|
||||
voice_id: str,
|
||||
request: Request,
|
||||
session: SessionDep,
|
||||
) -> JavaJSONResponse:
|
||||
user = require_normal(request)
|
||||
service = VoiceCloneService(session)
|
||||
await service.check_permission(voice_id, user)
|
||||
return ok(await service.create_audio_id(voice_id))
|
||||
|
||||
|
||||
@voiceclone_router.get("/voiceClone/play/{download_id}")
|
||||
async def play_voice_clone(download_id: str, session: SessionDep) -> Response:
|
||||
try:
|
||||
content = await VoiceCloneService(session).consume_audio(download_id)
|
||||
if content is None:
|
||||
return Response(status_code=404)
|
||||
return Response(
|
||||
content,
|
||||
media_type="audio/wav",
|
||||
headers={
|
||||
"Content-Length": str(len(content)),
|
||||
"Content-Disposition": "inline; filename=voice.wav",
|
||||
},
|
||||
)
|
||||
except Exception:
|
||||
return Response(status_code=500)
|
||||
|
||||
|
||||
@voiceclone_router.post("/voiceClone/cloneAudio")
|
||||
async def train_voice_clone(
|
||||
body: VoiceCloneTrainRequest,
|
||||
request: Request,
|
||||
session: SessionDep,
|
||||
) -> JavaJSONResponse:
|
||||
user = require_normal(request)
|
||||
service = VoiceCloneService(session)
|
||||
await service.check_permission(body.clone_id, user)
|
||||
await service.clone_audio(
|
||||
body.clone_id or "",
|
||||
accept_language=request.headers.get("Accept-Language"),
|
||||
)
|
||||
return ok()
|
||||
@@ -0,0 +1 @@
|
||||
"""Pydantic request and response schemas."""
|
||||
@@ -0,0 +1,238 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
from datetime import datetime
|
||||
from typing import Any
|
||||
|
||||
from pydantic import Field, field_validator
|
||||
from pydantic_core import PydanticCustomError
|
||||
|
||||
from app.schemas.common import JavaModel
|
||||
|
||||
|
||||
class AgentCreate(JavaModel):
|
||||
agent_name: str
|
||||
|
||||
@field_validator("agent_name", mode="before")
|
||||
@classmethod
|
||||
def require_non_blank_name(cls, value: Any) -> Any:
|
||||
if value is None or isinstance(value, str) and not value.strip():
|
||||
raise PydanticCustomError("java_not_blank", "智能体名称不能为空")
|
||||
return value
|
||||
|
||||
|
||||
class AgentMemory(JavaModel):
|
||||
summary_memory: str | None = None
|
||||
|
||||
|
||||
class ContextProvider(JavaModel):
|
||||
url: str | None = None
|
||||
headers: dict[str, Any] | None = None
|
||||
|
||||
|
||||
class FunctionInfo(JavaModel):
|
||||
plugin_id: str | None = None
|
||||
param_info: dict[str, Any] = Field(default_factory=dict)
|
||||
|
||||
@field_validator("param_info", mode="before")
|
||||
@classmethod
|
||||
def normalize_param_info(cls, value: Any) -> dict[str, Any]:
|
||||
if value is None or value == "":
|
||||
return {}
|
||||
if isinstance(value, str):
|
||||
parsed = json.loads(value)
|
||||
if not isinstance(parsed, dict):
|
||||
raise ValueError("paramInfo must be a JSON object")
|
||||
return {str(key): item for key, item in parsed.items()}
|
||||
if isinstance(value, dict):
|
||||
return {str(key): item for key, item in value.items() if key is not None}
|
||||
parsed = json.loads(json.dumps(value))
|
||||
if not isinstance(parsed, dict):
|
||||
raise ValueError("paramInfo must be an object")
|
||||
return {str(key): item for key, item in parsed.items()}
|
||||
|
||||
|
||||
class AgentUpdate(JavaModel):
|
||||
agent_code: str | None = None
|
||||
agent_name: str | None = None
|
||||
asr_model_id: str | None = None
|
||||
vad_model_id: str | None = None
|
||||
llm_model_id: str | None = None
|
||||
slm_model_id: str | None = None
|
||||
vllm_model_id: str | None = None
|
||||
tts_model_id: str | None = None
|
||||
tts_voice_id: str | None = None
|
||||
tts_language: str | None = None
|
||||
tts_volume: int | None = None
|
||||
tts_rate: int | None = None
|
||||
tts_pitch: int | None = None
|
||||
mem_model_id: str | None = None
|
||||
intent_model_id: str | None = None
|
||||
functions: list[FunctionInfo] | None = None
|
||||
system_prompt: str | None = None
|
||||
summary_memory: str | None = None
|
||||
chat_history_conf: int | None = None
|
||||
lang_code: str | None = None
|
||||
language: str | None = None
|
||||
sort: int | None = None
|
||||
context_providers: list[ContextProvider] | None = None
|
||||
correct_word_file_ids: list[str] | None = None
|
||||
tag_names: list[str] | None = None
|
||||
tag_ids: list[str] | None = None
|
||||
|
||||
|
||||
class AgentChatHistoryReport(JavaModel):
|
||||
mac_address: str
|
||||
session_id: str
|
||||
chat_type: int
|
||||
content: str
|
||||
audio_base64: str | None = None
|
||||
report_time: int | None = None
|
||||
|
||||
@field_validator("mac_address", "session_id", "content", mode="before")
|
||||
@classmethod
|
||||
def require_non_blank(cls, value: Any) -> Any:
|
||||
if value is None or isinstance(value, str) and not value.strip():
|
||||
raise PydanticCustomError("java_not_blank", "不能为空")
|
||||
return value
|
||||
|
||||
@field_validator("chat_type", mode="before")
|
||||
@classmethod
|
||||
def require_chat_type(cls, value: Any) -> Any:
|
||||
if value is None:
|
||||
raise PydanticCustomError("java_not_null", "不能为空")
|
||||
return value
|
||||
|
||||
|
||||
class AgentSnapshotPage(JavaModel):
|
||||
page: int | None = 1
|
||||
limit: int | None = 10
|
||||
max_version_no: int | None = None
|
||||
|
||||
def page_or_default(self) -> int:
|
||||
return self.page if self.page is not None and self.page >= 1 else 1
|
||||
|
||||
def limit_or_default(self) -> int:
|
||||
return self.limit if self.limit is not None and self.limit >= 1 else 10
|
||||
|
||||
|
||||
class AgentSnapshotRestore(JavaModel):
|
||||
current_state_token: str
|
||||
|
||||
@field_validator("current_state_token", mode="before")
|
||||
@classmethod
|
||||
def require_non_blank_token(cls, value: Any) -> Any:
|
||||
if value is None or isinstance(value, str) and not value.strip():
|
||||
raise PydanticCustomError("java_not_blank", "不能为空")
|
||||
return value
|
||||
|
||||
|
||||
class AgentSnapshotTag(JavaModel):
|
||||
id: str | None = None
|
||||
tag_name: str | None = None
|
||||
sort: int | None = None
|
||||
|
||||
|
||||
class AgentSnapshotData(JavaModel):
|
||||
agent_code: str | None = None
|
||||
agent_name: str | None = None
|
||||
asr_model_id: str | None = None
|
||||
vad_model_id: str | None = None
|
||||
llm_model_id: str | None = None
|
||||
slm_model_id: str | None = None
|
||||
vllm_model_id: str | None = None
|
||||
tts_model_id: str | None = None
|
||||
tts_voice_id: str | None = None
|
||||
tts_language: str | None = None
|
||||
tts_volume: int | None = None
|
||||
tts_rate: int | None = None
|
||||
tts_pitch: int | None = None
|
||||
mem_model_id: str | None = None
|
||||
intent_model_id: str | None = None
|
||||
chat_history_conf: int | None = None
|
||||
system_prompt: str | None = None
|
||||
summary_memory: str | None = None
|
||||
lang_code: str | None = None
|
||||
language: str | None = None
|
||||
sort: int | None = None
|
||||
functions: list[FunctionInfo] | None = None
|
||||
context_providers: list[ContextProvider] | None = None
|
||||
correct_word_file_ids: list[str] | None = None
|
||||
tag_names: list[str] | None = None
|
||||
tags: list[AgentSnapshotTag] | None = None
|
||||
|
||||
|
||||
class AgentTemplate(JavaModel):
|
||||
id: str | None = None
|
||||
agent_code: str | None = None
|
||||
agent_name: str | None = None
|
||||
asr_model_id: str | None = None
|
||||
vad_model_id: str | None = None
|
||||
llm_model_id: str | None = None
|
||||
vllm_model_id: str | None = None
|
||||
tts_model_id: str | None = None
|
||||
tts_voice_id: str | None = None
|
||||
tts_language: str | None = None
|
||||
tts_volume: int | None = None
|
||||
tts_rate: int | None = None
|
||||
tts_pitch: int | None = None
|
||||
mem_model_id: str | None = None
|
||||
intent_model_id: str | None = None
|
||||
chat_history_conf: int | None = None
|
||||
system_prompt: str | None = None
|
||||
summary_memory: str | None = None
|
||||
lang_code: str | None = None
|
||||
language: str | None = None
|
||||
sort: int | None = None
|
||||
creator: int | None = None
|
||||
created_at: datetime | None = None
|
||||
updater: int | None = None
|
||||
updated_at: datetime | None = None
|
||||
|
||||
|
||||
class AgentVoicePrintSave(JavaModel):
|
||||
agent_id: str | None = None
|
||||
audio_id: str | None = None
|
||||
source_name: str | None = None
|
||||
introduce: str | None = None
|
||||
|
||||
|
||||
class AgentVoicePrintUpdate(JavaModel):
|
||||
id: str | None = None
|
||||
audio_id: str | None = None
|
||||
source_name: str | None = None
|
||||
introduce: str | None = None
|
||||
|
||||
|
||||
class AgentTagAssignment(JavaModel):
|
||||
tag_ids: list[str] | None = None
|
||||
tag_names: list[str] | None = None
|
||||
|
||||
|
||||
SNAPSHOT_FIELD_ORDER = [
|
||||
"agentCode",
|
||||
"agentName",
|
||||
"asrModelId",
|
||||
"vadModelId",
|
||||
"llmModelId",
|
||||
"slmModelId",
|
||||
"vllmModelId",
|
||||
"ttsModelId",
|
||||
"ttsVoiceId",
|
||||
"ttsLanguage",
|
||||
"ttsVolume",
|
||||
"ttsRate",
|
||||
"ttsPitch",
|
||||
"memModelId",
|
||||
"intentModelId",
|
||||
"chatHistoryConf",
|
||||
"systemPrompt",
|
||||
"summaryMemory",
|
||||
"langCode",
|
||||
"language",
|
||||
"sort",
|
||||
"functions",
|
||||
"contextProviders",
|
||||
"correctWordFileIds",
|
||||
"tagNames",
|
||||
]
|
||||
@@ -0,0 +1,57 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from collections.abc import Callable
|
||||
from typing import Any, Generic, TypeVar
|
||||
|
||||
from pydantic import BaseModel, ConfigDict, Field
|
||||
|
||||
T = TypeVar("T")
|
||||
|
||||
|
||||
def to_camel(value: str) -> str:
|
||||
head, *tail = value.split("_")
|
||||
return head + "".join(part[:1].upper() + part[1:] for part in tail)
|
||||
|
||||
|
||||
class JavaModel(BaseModel):
|
||||
model_config = ConfigDict(
|
||||
alias_generator=to_camel,
|
||||
populate_by_name=True,
|
||||
extra="ignore",
|
||||
str_strip_whitespace=False,
|
||||
serialize_by_alias=True,
|
||||
)
|
||||
|
||||
|
||||
class PageData(JavaModel, Generic[T]):
|
||||
total: int
|
||||
list: list[T]
|
||||
|
||||
|
||||
class PageQuery(JavaModel):
|
||||
page: int = Field(default=1, ge=1)
|
||||
limit: int = Field(default=10, ge=1)
|
||||
order_field: str | list[str] | None = None
|
||||
order: str | None = None
|
||||
|
||||
|
||||
class DeleteIds(JavaModel):
|
||||
ids: list[str]
|
||||
|
||||
|
||||
def page_payload(rows: list[Any], total: int) -> dict[str, Any]:
|
||||
return {"total": int(total), "list": rows}
|
||||
|
||||
|
||||
def safe_order_by(
|
||||
requested: str | list[str] | None,
|
||||
*,
|
||||
allowed: set[str],
|
||||
default: str,
|
||||
transform: Callable[[str], str] | None = None,
|
||||
) -> list[str]:
|
||||
fields = [requested] if isinstance(requested, str) else list(requested or [])
|
||||
selected = [field for field in fields if field in allowed]
|
||||
if not selected:
|
||||
selected = [default]
|
||||
return [transform(field) if transform else field for field in selected]
|
||||
@@ -0,0 +1,25 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from pydantic import field_validator
|
||||
|
||||
from app.schemas.common import JavaModel
|
||||
|
||||
|
||||
def _not_blank(value: str) -> str:
|
||||
if not value or not value.strip():
|
||||
raise ValueError("must not be blank")
|
||||
return value
|
||||
|
||||
|
||||
class AgentModelsRequest(JavaModel):
|
||||
mac_address: str
|
||||
client_id: str
|
||||
selected_module: dict[str, str]
|
||||
|
||||
_validate_required = field_validator("mac_address", "client_id")(_not_blank)
|
||||
|
||||
|
||||
class CorrectWordsRequest(JavaModel):
|
||||
mac_address: str
|
||||
|
||||
_validate_required = field_validator("mac_address")(_not_blank)
|
||||
@@ -0,0 +1,9 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from app.schemas.common import JavaModel
|
||||
|
||||
|
||||
class CorrectWordFileBody(JavaModel):
|
||||
file_name: str | None = None
|
||||
content: list[str] | None = None
|
||||
file_size: int | None = None
|
||||
@@ -0,0 +1,145 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Any
|
||||
|
||||
from pydantic import AliasChoices, Field
|
||||
|
||||
from app.schemas.common import JavaModel
|
||||
|
||||
|
||||
class DeviceRegisterRequest(JavaModel):
|
||||
mac_address: str | None = None
|
||||
|
||||
|
||||
class DeviceUnbindRequest(JavaModel):
|
||||
device_id: str | None = None
|
||||
|
||||
|
||||
class DeviceUpdateRequest(JavaModel):
|
||||
auto_update: int | None = None
|
||||
alias: str | None = None
|
||||
|
||||
|
||||
class DeviceManualAddRequest(JavaModel):
|
||||
agent_id: str | None = None
|
||||
board: str | None = None
|
||||
app_version: str | None = None
|
||||
mac_address: str | None = None
|
||||
|
||||
|
||||
class DeviceToolCallRequest(JavaModel):
|
||||
name: str | None = None
|
||||
arguments: dict[str, Any] | None = None
|
||||
|
||||
|
||||
class DeviceAddressBookAliasRequest(JavaModel):
|
||||
mac_address: str | None = None
|
||||
target_mac: str | None = None
|
||||
alias: str | None = None
|
||||
|
||||
|
||||
class DeviceAddressBookPermissionRequest(JavaModel):
|
||||
mac_address: str | None = None
|
||||
target_mac: str | None = None
|
||||
has_permission: bool | None = None
|
||||
|
||||
|
||||
class ChipInfo(JavaModel):
|
||||
model: int | None = None
|
||||
cores: int | None = None
|
||||
revision: int | None = None
|
||||
features: int | None = None
|
||||
|
||||
|
||||
class ApplicationInfo(JavaModel):
|
||||
name: str | None = None
|
||||
version: str | None = None
|
||||
compile_time: str | None = Field(
|
||||
default=None,
|
||||
validation_alias=AliasChoices("compile_time", "compileTime"),
|
||||
serialization_alias="compile_time",
|
||||
)
|
||||
idf_version: str | None = Field(
|
||||
default=None,
|
||||
validation_alias=AliasChoices("idf_version", "idfVersion"),
|
||||
serialization_alias="idf_version",
|
||||
)
|
||||
elf_sha256: str | None = Field(
|
||||
default=None,
|
||||
validation_alias=AliasChoices("elf_sha256", "elfSha256"),
|
||||
serialization_alias="elf_sha256",
|
||||
)
|
||||
|
||||
|
||||
class PartitionInfo(JavaModel):
|
||||
label: str | None = None
|
||||
type: int | None = None
|
||||
subtype: int | None = None
|
||||
address: int | None = None
|
||||
size: int | None = None
|
||||
|
||||
|
||||
class OtaPartitionInfo(JavaModel):
|
||||
label: str | None = None
|
||||
|
||||
|
||||
class BoardInfo(JavaModel):
|
||||
type: str | None = None
|
||||
ssid: str | None = None
|
||||
rssi: int | None = None
|
||||
channel: int | None = None
|
||||
ip: str | None = None
|
||||
mac: str | None = None
|
||||
|
||||
|
||||
class DeviceReportRequest(JavaModel):
|
||||
version: int | None = None
|
||||
flash_size: int | None = Field(
|
||||
default=None,
|
||||
validation_alias=AliasChoices("flash_size", "flashSize"),
|
||||
serialization_alias="flash_size",
|
||||
)
|
||||
minimum_free_heap_size: int | None = Field(
|
||||
default=None,
|
||||
validation_alias=AliasChoices("minimum_free_heap_size", "minimumFreeHeapSize"),
|
||||
serialization_alias="minimum_free_heap_size",
|
||||
)
|
||||
mac_address: str | None = Field(
|
||||
default=None,
|
||||
validation_alias=AliasChoices("mac_address", "macAddress"),
|
||||
serialization_alias="mac_address",
|
||||
)
|
||||
uuid: str | None = None
|
||||
chip_model_name: str | None = Field(
|
||||
default=None,
|
||||
validation_alias=AliasChoices("chip_model_name", "chipModelName"),
|
||||
serialization_alias="chip_model_name",
|
||||
)
|
||||
chip_info: ChipInfo | None = Field(
|
||||
default=None,
|
||||
validation_alias=AliasChoices("chip_info", "chipInfo"),
|
||||
serialization_alias="chip_info",
|
||||
)
|
||||
application: ApplicationInfo | None = None
|
||||
partition_table: list[PartitionInfo] | None = Field(
|
||||
default=None,
|
||||
validation_alias=AliasChoices("partition_table", "partitionTable"),
|
||||
serialization_alias="partition_table",
|
||||
)
|
||||
ota: OtaPartitionInfo | None = None
|
||||
board: BoardInfo | None = None
|
||||
|
||||
|
||||
class OtaRecord(JavaModel):
|
||||
id: str | None = None
|
||||
firmware_name: str | None = None
|
||||
type: str | None = None
|
||||
version: str | None = None
|
||||
size: int | None = None
|
||||
remark: str | None = None
|
||||
firmware_path: str | None = None
|
||||
sort: int | None = None
|
||||
updater: int | None = None
|
||||
update_date: str | None = None
|
||||
creator: int | None = None
|
||||
create_date: str | None = None
|
||||
@@ -0,0 +1,53 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from datetime import datetime
|
||||
from typing import Any
|
||||
|
||||
from pydantic import AliasChoices, Field
|
||||
|
||||
from app.schemas.common import JavaModel
|
||||
|
||||
|
||||
class KnowledgeBaseBody(JavaModel):
|
||||
id: str | None = None
|
||||
dataset_id: str | None = None
|
||||
rag_model_id: str | None = None
|
||||
name: str | None = None
|
||||
avatar: str | None = None
|
||||
description: str | None = None
|
||||
embedding_model: str | None = None
|
||||
permission: str | None = None
|
||||
chunk_method: str | None = None
|
||||
parser_config: str | None = None
|
||||
chunk_count: int | None = None
|
||||
token_num: int | None = None
|
||||
status: int | None = None
|
||||
creator: int | None = None
|
||||
created_at: datetime | None = None
|
||||
updater: int | None = None
|
||||
updated_at: datetime | None = None
|
||||
document_count: int | None = None
|
||||
error_message: str | None = None
|
||||
|
||||
|
||||
class DocumentBatchBody(JavaModel):
|
||||
ids: list[str] | None = Field(
|
||||
default=None,
|
||||
validation_alias=AliasChoices("ids", "document_ids"),
|
||||
)
|
||||
|
||||
|
||||
class RetrievalBody(JavaModel):
|
||||
dataset_ids: list[str] | None = None
|
||||
document_ids: list[str] | None = None
|
||||
question: str | None = None
|
||||
page: int | None = None
|
||||
page_size: int | None = None
|
||||
similarity_threshold: float | None = None
|
||||
vector_similarity_weight: float | None = None
|
||||
top_k: int | None = None
|
||||
rerank_id: str | None = None
|
||||
highlight: bool | None = None
|
||||
keyword: bool | None = None
|
||||
cross_languages: list[str] | None = None
|
||||
metadata_condition: dict[str, Any] | None = None
|
||||
@@ -0,0 +1,25 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Any
|
||||
|
||||
from app.schemas.common import JavaModel
|
||||
|
||||
|
||||
class ModelConfigBody(JavaModel):
|
||||
id: str | None = None
|
||||
model_code: str | None = None
|
||||
model_name: str | None = None
|
||||
is_default: int | None = None
|
||||
is_enabled: int | None = None
|
||||
config_json: dict[str, Any] | None = None
|
||||
doc_link: str | None = None
|
||||
remark: str | None = None
|
||||
sort: int | None = None
|
||||
|
||||
class ModelProviderBody(JavaModel):
|
||||
id: str | None = None
|
||||
model_type: str | None = None
|
||||
provider_code: str | None = None
|
||||
name: str | None = None
|
||||
fields: str | None = None
|
||||
sort: int | None = None
|
||||
@@ -0,0 +1,59 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Any
|
||||
|
||||
from app.schemas.common import JavaModel
|
||||
|
||||
|
||||
class LoginRequest(JavaModel):
|
||||
# LoginController does not use @Valid; null/blank values reach its service logic.
|
||||
username: str | None = None
|
||||
password: str | None = None
|
||||
mobile_captcha: str | None = None
|
||||
captcha_id: str | None = None
|
||||
|
||||
|
||||
class SmsVerificationRequest(JavaModel):
|
||||
# smsVerification likewise omits @Valid in the Java controller.
|
||||
phone: str | None = None
|
||||
captcha: str | None = None
|
||||
captcha_id: str | None = None
|
||||
|
||||
|
||||
class PasswordChangeRequest(JavaModel):
|
||||
password: str | None = None
|
||||
new_password: str | None = None
|
||||
|
||||
|
||||
class RetrievePasswordRequest(JavaModel):
|
||||
phone: str | None = None
|
||||
code: str | None = None
|
||||
password: str | None = None
|
||||
captcha_id: str | None = None
|
||||
|
||||
|
||||
class TokenData(JavaModel):
|
||||
token: str
|
||||
expire: int
|
||||
client_hash: str | None
|
||||
|
||||
|
||||
class UserDetailData(JavaModel):
|
||||
id: int
|
||||
username: str
|
||||
super_admin: int
|
||||
token: str
|
||||
status: int
|
||||
|
||||
|
||||
class PublicConfigData(JavaModel):
|
||||
enable_mobile_register: bool
|
||||
version: str
|
||||
year: str
|
||||
allow_user_register: bool
|
||||
mobile_area_list: list[dict[str, Any]]
|
||||
beian_icp_num: str | None
|
||||
beian_ga_num: str | None
|
||||
name: str | None
|
||||
sm2_public_key: str
|
||||
system_web_menu: Any | None = None
|
||||
@@ -0,0 +1,45 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from pydantic import field_validator
|
||||
|
||||
from app.schemas.common import JavaModel
|
||||
|
||||
|
||||
def _not_blank(value: str) -> str:
|
||||
if not value or not value.strip():
|
||||
raise ValueError("must not be blank")
|
||||
return value
|
||||
|
||||
|
||||
class SysParamPayload(JavaModel):
|
||||
id: int | None = None
|
||||
param_code: str | None = None
|
||||
param_value: str | None = None
|
||||
value_type: str | None = None
|
||||
remark: str | None = None
|
||||
|
||||
|
||||
class DictTypePayload(JavaModel):
|
||||
# Controller calls ValidatorUtils without the DTO's custom groups, so these constraints don't execute in Java.
|
||||
id: int | None = None
|
||||
dict_type: str | None = None
|
||||
dict_name: str | None = None
|
||||
remark: str | None = None
|
||||
sort: int | None = None
|
||||
|
||||
|
||||
class DictDataPayload(JavaModel):
|
||||
# See DictTypePayload: Add/Update/DefaultGroup annotations are skipped by the Java controller.
|
||||
id: int | None = None
|
||||
dict_type_id: int | None = None
|
||||
dict_label: str | None = None
|
||||
dict_value: str | None = None
|
||||
remark: str | None = None
|
||||
sort: int | None = None
|
||||
|
||||
|
||||
class EmitServerActionRequest(JavaModel):
|
||||
target_ws: str
|
||||
action: str | None
|
||||
|
||||
_validate_target = field_validator("target_ws")(_not_blank)
|
||||
@@ -0,0 +1,15 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from app.schemas.common import JavaModel
|
||||
|
||||
|
||||
class TimbreBody(JavaModel):
|
||||
languages: str | None = None
|
||||
name: str | None = None
|
||||
remark: str | None = None
|
||||
reference_audio: str | None = None
|
||||
reference_text: str | None = None
|
||||
sort: int | None = 0
|
||||
tts_model_id: str | None = None
|
||||
tts_voice: str | None = None
|
||||
voice_demo: str | None = None
|
||||
@@ -0,0 +1,19 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from app.schemas.common import JavaModel
|
||||
|
||||
|
||||
class VoiceResourceCreateRequest(JavaModel):
|
||||
model_id: str | None = None
|
||||
voice_ids: list[str] | None = None
|
||||
user_id: int | None = None
|
||||
languages: str | None = None
|
||||
|
||||
|
||||
class VoiceCloneRenameRequest(JavaModel):
|
||||
id: str | None = None
|
||||
name: str | None = None
|
||||
|
||||
|
||||
class VoiceCloneTrainRequest(JavaModel):
|
||||
clone_id: str | None = None
|
||||
@@ -0,0 +1 @@
|
||||
"""Business services."""
|
||||
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,521 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import base64
|
||||
import hashlib
|
||||
import json
|
||||
import logging
|
||||
import math
|
||||
import urllib.parse
|
||||
from copy import deepcopy
|
||||
from typing import Any, cast
|
||||
|
||||
from cryptography.hazmat.primitives.ciphers import Cipher, algorithms, modes
|
||||
from redis.asyncio import Redis
|
||||
|
||||
from app.core.errors import AppError
|
||||
from app.core.redis import JavaRedisCodec, get_redis
|
||||
from app.repositories.config import ConfigRepository
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
class ConfigService:
|
||||
def __init__(self, repository: ConfigRepository, *, redis: Redis | None = None):
|
||||
self.repository = repository
|
||||
self.redis = redis or get_redis()
|
||||
|
||||
async def get_config(self, *, use_cache: bool) -> dict[str, Any]:
|
||||
if use_cache:
|
||||
cached = JavaRedisCodec.decode(await cast(Any, self.redis.get)("server:config"))
|
||||
if isinstance(cached, dict):
|
||||
return cast(dict[str, Any], cached)
|
||||
result = self._build_base_config(await self.repository.list_params())
|
||||
template = await self.repository.get_default_template()
|
||||
if template is None:
|
||||
raise AppError(10183)
|
||||
await self._build_module_config(
|
||||
result=result,
|
||||
assistant_name=None,
|
||||
prompt=None,
|
||||
summary_memory=None,
|
||||
voice=None,
|
||||
reference_audio=None,
|
||||
reference_text=None,
|
||||
language=None,
|
||||
tts_volume=None,
|
||||
tts_rate=None,
|
||||
tts_pitch=None,
|
||||
vad_model_id=self._string(template.get("vad_model_id")),
|
||||
asr_model_id=self._string(template.get("asr_model_id")),
|
||||
llm_model_id=None,
|
||||
vllm_model_id=None,
|
||||
slm_model_id=None,
|
||||
tts_model_id=None,
|
||||
mem_model_id=None,
|
||||
intent_model_id=None,
|
||||
rag_model_id=None,
|
||||
)
|
||||
await cast(Any, self.redis.set)("server:config", JavaRedisCodec.encode(result), ex=24 * 60 * 60)
|
||||
return result
|
||||
|
||||
async def get_agent_models(
|
||||
self,
|
||||
mac_address: str,
|
||||
selected_module: dict[str, str],
|
||||
) -> dict[str, Any]:
|
||||
temporary_key = f"tmp_register_mac:{mac_address}"
|
||||
temporary = JavaRedisCodec.decode(await cast(Any, self.redis.get)(temporary_key))
|
||||
if temporary == "true":
|
||||
await cast(Any, self.redis.delete)(temporary_key)
|
||||
return await self.get_config(use_cache=True)
|
||||
|
||||
device = await self.repository.get_device_by_mac(mac_address)
|
||||
if device is None:
|
||||
safe_address = mac_address.replace(":", "_").lower()
|
||||
activation = JavaRedisCodec.decode(
|
||||
await cast(Any, self.redis.get)(f"ota:activation:data:{safe_address}")
|
||||
)
|
||||
if isinstance(activation, dict) and activation.get("activation_code"):
|
||||
raise AppError(10042, params=(str(activation["activation_code"]),))
|
||||
raise AppError(10041)
|
||||
|
||||
agent_id = self._string(device.get("agent_id"))
|
||||
agent = await self.repository.get_agent(agent_id or "") if agent_id else None
|
||||
if agent is None:
|
||||
raise AppError(10053)
|
||||
|
||||
voice: str | None = None
|
||||
reference_audio: str | None = None
|
||||
reference_text: str | None = None
|
||||
language: str | None = None
|
||||
voice_id = self._string(agent.get("tts_voice_id"))
|
||||
timbre = await self._timbre(voice_id) if voice_id else None
|
||||
if timbre is not None:
|
||||
voice = self._string(timbre.get("tts_voice"))
|
||||
reference_audio = self._string(timbre.get("reference_audio"))
|
||||
reference_text = self._string(timbre.get("reference_text"))
|
||||
chosen_language = self._string(agent.get("tts_language"))
|
||||
if chosen_language and chosen_language.strip():
|
||||
language = chosen_language
|
||||
else:
|
||||
languages = self._string(timbre.get("languages"))
|
||||
if languages and languages.strip():
|
||||
language = languages.split("、", 1)[0].strip()
|
||||
elif voice_id:
|
||||
clone = await self.repository.get_voice_clone(voice_id)
|
||||
if clone is not None:
|
||||
voice = self._string(clone.get("voice_id"))
|
||||
chosen_language = self._string(agent.get("tts_language"))
|
||||
language = chosen_language if chosen_language and chosen_language.strip() else "普通话"
|
||||
|
||||
result: dict[str, Any] = {
|
||||
"device_max_output_size": await self._param("device_max_output_size", from_cache=True)
|
||||
}
|
||||
memory_model = self._string(agent.get("mem_model_id"))
|
||||
chat_history = agent.get("chat_history_conf")
|
||||
if memory_model == "Memory_nomem":
|
||||
chat_history = 0
|
||||
elif memory_model is not None and memory_model != "Memory_nomem" and chat_history is None:
|
||||
chat_history = 2
|
||||
result["chat_history_conf"] = chat_history
|
||||
|
||||
vad_model_id = self._string(agent.get("vad_model_id"))
|
||||
asr_model_id = self._string(agent.get("asr_model_id"))
|
||||
if selected_module.get("VAD") == vad_model_id:
|
||||
vad_model_id = None
|
||||
if selected_module.get("ASR") == asr_model_id:
|
||||
asr_model_id = None
|
||||
|
||||
if self._string(agent.get("intent_model_id")) != "Intent_nointent":
|
||||
plugins = await self._plugins(str(agent["id"]))
|
||||
if plugins:
|
||||
result["plugins"] = plugins
|
||||
|
||||
mcp_endpoint = await self._mcp_address(str(agent["id"]))
|
||||
if mcp_endpoint and mcp_endpoint.startswith("ws"):
|
||||
result["mcp_endpoint"] = mcp_endpoint.replace("/mcp/", "/call/")
|
||||
|
||||
context_providers = self._json_value(await self.repository.get_context_providers(str(agent["id"])))
|
||||
if isinstance(context_providers, list) and context_providers:
|
||||
result["context_providers"] = context_providers
|
||||
|
||||
await self._add_voiceprint(str(agent["id"]), result)
|
||||
await self._build_module_config(
|
||||
result=result,
|
||||
assistant_name=self._string(agent.get("agent_name")),
|
||||
prompt=self._string(agent.get("system_prompt")),
|
||||
summary_memory=self._string(agent.get("summary_memory")),
|
||||
voice=voice,
|
||||
reference_audio=reference_audio,
|
||||
reference_text=reference_text,
|
||||
language=language,
|
||||
tts_volume=self._integer(agent.get("tts_volume")),
|
||||
tts_rate=self._integer(agent.get("tts_rate")),
|
||||
tts_pitch=self._integer(agent.get("tts_pitch")),
|
||||
vad_model_id=vad_model_id,
|
||||
asr_model_id=asr_model_id,
|
||||
llm_model_id=self._string(agent.get("llm_model_id")),
|
||||
vllm_model_id=self._string(agent.get("vllm_model_id")),
|
||||
slm_model_id=self._string(agent.get("slm_model_id")),
|
||||
tts_model_id=self._string(agent.get("tts_model_id")),
|
||||
mem_model_id=memory_model,
|
||||
intent_model_id=self._string(agent.get("intent_model_id")),
|
||||
rag_model_id=None,
|
||||
)
|
||||
return result
|
||||
|
||||
async def get_correct_words(self, mac_address: str) -> list[str]:
|
||||
device = await self.repository.get_device_by_mac(mac_address)
|
||||
if device is None or device.get("agent_id") is None:
|
||||
return []
|
||||
rows = await self.repository.get_correct_word_items(str(device["agent_id"]))
|
||||
return [
|
||||
f"{self._java_string(row.get('source_word'))}|{self._java_string(row.get('target_word'))}"
|
||||
for row in rows
|
||||
]
|
||||
|
||||
@staticmethod
|
||||
def _build_base_config(rows: list[dict[str, Any]]) -> dict[str, Any]:
|
||||
config: dict[str, Any] = {}
|
||||
for row in rows:
|
||||
code = str(row.get("param_code") or "")
|
||||
keys = code.split(".")
|
||||
current = config
|
||||
for key in keys[:-1]:
|
||||
if key not in current:
|
||||
current[key] = {}
|
||||
nested = current[key]
|
||||
if not isinstance(nested, dict):
|
||||
raise TypeError(f"configuration path {code} collides with scalar key {key}")
|
||||
current = nested
|
||||
value = str(row.get("param_value") or "")
|
||||
value_type = str(row.get("value_type") or "string").lower()
|
||||
current[keys[-1]] = ConfigService._typed_param(value, value_type)
|
||||
return config
|
||||
|
||||
@staticmethod
|
||||
def _typed_param(value: str, value_type: str) -> Any:
|
||||
if value_type == "number":
|
||||
try:
|
||||
number = float(value)
|
||||
# Java's implementation returns an Integer only when the double
|
||||
# equals its narrowing conversion to a signed 32-bit int.
|
||||
if math.isnan(number):
|
||||
narrowed = 0
|
||||
elif number >= 2**31 - 1:
|
||||
narrowed = 2**31 - 1
|
||||
elif number <= -(2**31):
|
||||
narrowed = -(2**31)
|
||||
else:
|
||||
narrowed = int(number)
|
||||
return narrowed if number == narrowed else number
|
||||
except ValueError:
|
||||
return value
|
||||
if value_type == "boolean":
|
||||
return value.lower() == "true"
|
||||
if value_type == "array":
|
||||
return [item.strip() for item in value.split(";") if item.strip()]
|
||||
if value_type == "json":
|
||||
try:
|
||||
return json.loads(value)
|
||||
except json.JSONDecodeError:
|
||||
return value
|
||||
return value
|
||||
|
||||
async def _build_module_config(
|
||||
self,
|
||||
*,
|
||||
result: dict[str, Any],
|
||||
assistant_name: str | None,
|
||||
prompt: str | None,
|
||||
summary_memory: str | None,
|
||||
voice: str | None,
|
||||
reference_audio: str | None,
|
||||
reference_text: str | None,
|
||||
language: str | None,
|
||||
tts_volume: int | None,
|
||||
tts_rate: int | None,
|
||||
tts_pitch: int | None,
|
||||
vad_model_id: str | None,
|
||||
asr_model_id: str | None,
|
||||
llm_model_id: str | None,
|
||||
vllm_model_id: str | None,
|
||||
slm_model_id: str | None,
|
||||
tts_model_id: str | None,
|
||||
mem_model_id: str | None,
|
||||
intent_model_id: str | None,
|
||||
rag_model_id: str | None,
|
||||
) -> None:
|
||||
selected: dict[str, str] = {}
|
||||
model_types = ("VAD", "ASR", "TTS", "Memory", "Intent", "LLM", "VLLM", "SLM", "RAG")
|
||||
model_ids = (
|
||||
vad_model_id,
|
||||
asr_model_id,
|
||||
tts_model_id,
|
||||
mem_model_id,
|
||||
intent_model_id,
|
||||
llm_model_id,
|
||||
vllm_model_id,
|
||||
slm_model_id,
|
||||
rag_model_id,
|
||||
)
|
||||
intent_llm_id: str | None = None
|
||||
memory_llm_id: str | None = None
|
||||
for model_type, model_id in zip(model_types, model_ids, strict=True):
|
||||
if model_id is None:
|
||||
continue
|
||||
model = await self._model(model_id)
|
||||
if model is None:
|
||||
continue
|
||||
configuration = self._json_value(model.get("config_json"))
|
||||
type_config: dict[str, Any] = {}
|
||||
if isinstance(configuration, dict):
|
||||
configuration = deepcopy(configuration)
|
||||
type_config[str(model["id"])] = configuration
|
||||
if model_type == "TTS":
|
||||
optional_values = {
|
||||
"private_voice": voice,
|
||||
"ref_audio": reference_audio,
|
||||
"ref_text": reference_text,
|
||||
"language": language,
|
||||
"ttsVolume": tts_volume,
|
||||
"ttsRate": tts_rate,
|
||||
"ttsPitch": tts_pitch,
|
||||
}
|
||||
configuration.update({key: value for key, value in optional_values.items() if value is not None})
|
||||
if configuration.get("type") == "huoshan_double_stream" and voice and voice.startswith("S_"):
|
||||
configuration["resource_id"] = "seed-icl-1.0"
|
||||
elif model_type == "Intent":
|
||||
if configuration.get("type") == "intent_llm":
|
||||
intent_llm_id = self._string(configuration.get("llm"))
|
||||
if intent_llm_id == llm_model_id:
|
||||
intent_llm_id = None
|
||||
functions = configuration.get("functions")
|
||||
if isinstance(functions, str) and functions.strip():
|
||||
configuration["functions"] = functions.split(";")
|
||||
elif model_type == "Memory" and configuration.get("type") == "mem_local_short":
|
||||
memory_llm_id = self._string(configuration.get("llm"))
|
||||
if memory_llm_id == llm_model_id:
|
||||
memory_llm_id = None
|
||||
elif model_type == "LLM":
|
||||
for extra_id in (intent_llm_id, memory_llm_id):
|
||||
if extra_id and extra_id not in type_config:
|
||||
extra = await self._model(extra_id)
|
||||
if extra is not None:
|
||||
type_config[str(extra["id"])] = deepcopy(self._json_value(extra.get("config_json")))
|
||||
if slm_model_id and slm_model_id != llm_model_id and slm_model_id not in type_config:
|
||||
small = await self._model(slm_model_id)
|
||||
small_config = None if small is None else self._json_value(small.get("config_json"))
|
||||
if small is not None and small_config is not None:
|
||||
type_config[str(small["id"])] = deepcopy(small_config)
|
||||
result[model_type] = type_config
|
||||
selected[model_type] = str(model["id"])
|
||||
result["selected_module"] = selected
|
||||
if prompt and prompt.strip():
|
||||
replacement = assistant_name if assistant_name and assistant_name.strip() else "小智"
|
||||
prompt = prompt.replace("{{assistant_name}}", replacement)
|
||||
result["prompt"] = prompt
|
||||
result["summaryMemory"] = summary_memory
|
||||
|
||||
async def _plugins(self, agent_id: str) -> dict[str, Any]:
|
||||
mappings = await self.repository.get_plugin_mappings(agent_id)
|
||||
result: dict[str, Any] = {}
|
||||
knowledge_groups: dict[str, list[dict[str, Any]]] = {}
|
||||
knowledge_models: dict[str, dict[str, Any]] = {}
|
||||
for mapping in mappings:
|
||||
provider_code = self._string(mapping.get("provider_code"))
|
||||
if provider_code and provider_code.strip():
|
||||
value = mapping.get("param_info")
|
||||
result[provider_code] = (
|
||||
json.dumps(value, ensure_ascii=False, separators=(",", ":")) if isinstance(value, dict) else value
|
||||
)
|
||||
# Java removes knowledge mappings by iterating the original list backwards, which reverses dataset order.
|
||||
for mapping in reversed(mappings):
|
||||
provider_code = self._string(mapping.get("provider_code"))
|
||||
if provider_code and provider_code.strip():
|
||||
continue
|
||||
dataset = await self.repository.get_dataset(str(mapping["plugin_id"]))
|
||||
if dataset is None or dataset.get("rag_model_id") is None:
|
||||
continue
|
||||
model = await self._model(str(dataset["rag_model_id"]))
|
||||
if model is None or not model.get("model_code"):
|
||||
continue
|
||||
code = str(model["model_code"])
|
||||
knowledge_groups.setdefault(code, []).append(dataset)
|
||||
knowledge_models[code] = model
|
||||
for code, datasets in knowledge_groups.items():
|
||||
model_config = self._json_value(knowledge_models[code].get("config_json"))
|
||||
if not isinstance(model_config, dict):
|
||||
continue
|
||||
names = ",".join(self._java_string(dataset.get("name")) for dataset in datasets)
|
||||
descriptions = ",".join(
|
||||
self._java_string(dataset.get("description")) for dataset in datasets
|
||||
)
|
||||
params = {
|
||||
"base_url": model_config.get("base_url"),
|
||||
"api_key": model_config.get("api_key"),
|
||||
"dataset_ids": [dataset.get("dataset_id") for dataset in datasets],
|
||||
"description": (
|
||||
f"如果用户询问与【{names}】涵盖的主体范围相关内容时应调用本方法,"
|
||||
f"用于查询:{descriptions}"
|
||||
),
|
||||
}
|
||||
result[f"search_from_{code}"] = json.dumps(params, ensure_ascii=False, separators=(",", ":"))
|
||||
return result
|
||||
|
||||
async def _mcp_address(self, agent_id: str) -> str | None:
|
||||
endpoint = await self._param("server.mcp_endpoint", from_cache=True)
|
||||
if endpoint is None or not endpoint.strip() or endpoint == "null":
|
||||
return None
|
||||
parsed = urllib.parse.urlsplit(endpoint)
|
||||
query = parsed.query
|
||||
marker_index = query.find("key=")
|
||||
key = query[marker_index + len("key=") :]
|
||||
scheme = "wss" if parsed.scheme == "https" else "ws"
|
||||
path = parsed.path
|
||||
prefix_path = path[: path.rfind("/")] if "/" in path else ""
|
||||
prefix = urllib.parse.urlunsplit((scheme, parsed.netloc, prefix_path, "", ""))
|
||||
token = self._aes_encrypt(
|
||||
key,
|
||||
json.dumps(
|
||||
{"agentId": hashlib.md5(agent_id.encode(), usedforsecurity=False).hexdigest()},
|
||||
ensure_ascii=False,
|
||||
separators=(", ", ": "),
|
||||
),
|
||||
)
|
||||
return f"{prefix}/mcp/?token={urllib.parse.quote_plus(token)}"
|
||||
|
||||
@staticmethod
|
||||
def _aes_encrypt(key: str, plaintext: str) -> str:
|
||||
key_bytes = key.encode()
|
||||
if len(key_bytes) not in {16, 24, 32}:
|
||||
key_bytes = (key_bytes + bytes(32))[:32]
|
||||
block_size = 16
|
||||
padding_length = block_size - len(plaintext.encode()) % block_size
|
||||
padded = plaintext.encode() + bytes([padding_length]) * padding_length
|
||||
# Java's published MCP token format is AES/ECB/PKCS5Padding; changing modes breaks existing servers.
|
||||
encryptor = Cipher(algorithms.AES(key_bytes), modes.ECB()).encryptor() # noqa: S305
|
||||
encrypted = encryptor.update(padded) + encryptor.finalize()
|
||||
return base64.b64encode(encrypted).decode("ascii")
|
||||
|
||||
async def _add_voiceprint(self, agent_id: str, result: dict[str, Any]) -> None:
|
||||
try:
|
||||
url = await self._param("server.voice_print", from_cache=True)
|
||||
if url is None or not url.strip() or url == "null":
|
||||
return
|
||||
rows = await self.repository.get_voiceprints(agent_id)
|
||||
if not rows:
|
||||
return
|
||||
speakers = [
|
||||
(
|
||||
f"{self._java_string(row.get('id'))},"
|
||||
f"{self._java_string(row.get('source_name'))},{row.get('introduce') or ''}"
|
||||
)
|
||||
for row in rows
|
||||
]
|
||||
threshold_value = await self._param("server.voiceprint_similarity_threshold", from_cache=True)
|
||||
try:
|
||||
threshold = (
|
||||
float(threshold_value)
|
||||
if threshold_value is not None and threshold_value not in ("", "null")
|
||||
else 0.4
|
||||
)
|
||||
except ValueError:
|
||||
threshold = 0.4
|
||||
result["voiceprint"] = {"url": url, "speakers": speakers, "similarity_threshold": threshold}
|
||||
except Exception:
|
||||
logger.warning("Voiceprint configuration lookup failed", exc_info=True)
|
||||
|
||||
async def _param(self, code: str, *, from_cache: bool) -> str | None:
|
||||
if from_cache:
|
||||
cached = JavaRedisCodec.decode(await cast(Any, self.redis.hget)("sys:params", code))
|
||||
if cached is not None:
|
||||
return str(cached)
|
||||
value = await self.repository.get_param_value(code)
|
||||
if value is not None and from_cache:
|
||||
await cast(Any, self.redis.hset)("sys:params", code, JavaRedisCodec.encode(value))
|
||||
await cast(Any, self.redis.expire)("sys:params", 24 * 60 * 60)
|
||||
return value
|
||||
|
||||
async def _model(self, model_id: str) -> dict[str, Any] | None:
|
||||
key = f"model:data:{model_id}"
|
||||
cached = JavaRedisCodec.decode(await cast(Any, self.redis.get)(key))
|
||||
if isinstance(cached, dict):
|
||||
return self._normalize_cached(cast(dict[str, Any], cached))
|
||||
model = await self.repository.get_model(model_id)
|
||||
if model is not None:
|
||||
raw_configuration = model.get("config_json")
|
||||
if isinstance(raw_configuration, str):
|
||||
parsed_configuration = json.loads(raw_configuration)
|
||||
if parsed_configuration is not None and not isinstance(parsed_configuration, dict):
|
||||
raise TypeError("ModelConfigEntity.configJson must be a JSON object")
|
||||
model["config_json"] = parsed_configuration
|
||||
await cast(Any, self.redis.set)(
|
||||
key,
|
||||
JavaRedisCodec.encode(
|
||||
model,
|
||||
java_type="xiaozhi.modules.model.entity.ModelConfigEntity",
|
||||
field_java_types={
|
||||
"configJson": "cn.hutool.json.JSONObject",
|
||||
"creator": "java.lang.Long",
|
||||
"updater": "java.lang.Long",
|
||||
},
|
||||
),
|
||||
ex=24 * 60 * 60,
|
||||
)
|
||||
return model
|
||||
|
||||
async def _timbre(self, timbre_id: str) -> dict[str, Any] | None:
|
||||
key = f"timbre:details:{timbre_id}"
|
||||
cached = JavaRedisCodec.decode(await cast(Any, self.redis.get)(key))
|
||||
if isinstance(cached, dict):
|
||||
return self._normalize_cached(cast(dict[str, Any], cached))
|
||||
timbre = await self.repository.get_timbre(timbre_id)
|
||||
if timbre is not None:
|
||||
await cast(Any, self.redis.set)(
|
||||
key,
|
||||
JavaRedisCodec.encode(
|
||||
timbre,
|
||||
java_type="xiaozhi.modules.timbre.vo.TimbreDetailsVO",
|
||||
field_java_types={"sort": "java.lang.Long"},
|
||||
),
|
||||
ex=24 * 60 * 60,
|
||||
)
|
||||
return timbre
|
||||
|
||||
@staticmethod
|
||||
def _normalize_cached(value: dict[str, Any]) -> dict[str, Any]:
|
||||
aliases = {
|
||||
"modelType": "model_type",
|
||||
"modelCode": "model_code",
|
||||
"modelName": "model_name",
|
||||
"configJson": "config_json",
|
||||
"ttsVoice": "tts_voice",
|
||||
"referenceAudio": "reference_audio",
|
||||
"referenceText": "reference_text",
|
||||
"ttsModelId": "tts_model_id",
|
||||
}
|
||||
return {aliases.get(key, key): item for key, item in value.items() if key != "@class"}
|
||||
|
||||
@staticmethod
|
||||
def _json_value(value: Any) -> Any:
|
||||
if isinstance(value, bytes):
|
||||
value = value.decode()
|
||||
if isinstance(value, str):
|
||||
try:
|
||||
return json.loads(value)
|
||||
except json.JSONDecodeError:
|
||||
return value
|
||||
return value
|
||||
|
||||
@staticmethod
|
||||
def _string(value: Any) -> str | None:
|
||||
return None if value is None else str(value)
|
||||
|
||||
@staticmethod
|
||||
def _integer(value: Any) -> int | None:
|
||||
return None if value is None else int(value)
|
||||
|
||||
@staticmethod
|
||||
def _java_string(value: Any) -> str:
|
||||
return "null" if value is None else str(value)
|
||||
@@ -0,0 +1,133 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import uuid
|
||||
from typing import Any
|
||||
|
||||
from app.core.errors import AppError
|
||||
from app.core.security import AuthUser, shanghai_now_naive
|
||||
from app.repositories.correctword import CorrectWordRepository
|
||||
from app.schemas.correctword import CorrectWordFileBody
|
||||
|
||||
|
||||
def _parse_lines(lines: list[str]) -> list[tuple[str, str]]:
|
||||
result: list[tuple[str, str]] = []
|
||||
for raw in lines:
|
||||
line = raw.strip()
|
||||
if not line or "|" not in line:
|
||||
continue
|
||||
source, target = line.split("|", 1)
|
||||
if source.strip() and target.strip():
|
||||
result.append((source.strip(), target.strip()))
|
||||
return result
|
||||
|
||||
|
||||
def _content_lines(value: str | None) -> list[str]:
|
||||
if value is None:
|
||||
return []
|
||||
# Java String.split keeps one empty element for the empty source string,
|
||||
# while still discarding trailing empty elements for non-empty strings.
|
||||
if value == "":
|
||||
return [""]
|
||||
lines = value.split("\n")
|
||||
while lines and lines[-1] == "":
|
||||
lines.pop()
|
||||
return lines
|
||||
|
||||
|
||||
def file_vo(row: dict[str, Any]) -> dict[str, Any]:
|
||||
return {
|
||||
"id": row.get("id"),
|
||||
"fileName": row.get("file_name"),
|
||||
"wordCount": row.get("word_count"),
|
||||
"content": _content_lines(row.get("content")),
|
||||
"createdAt": row.get("created_at"),
|
||||
"updatedAt": row.get("updated_at"),
|
||||
}
|
||||
|
||||
|
||||
class CorrectWordService:
|
||||
def __init__(self, repository: CorrectWordRepository):
|
||||
self.repository = repository
|
||||
|
||||
@staticmethod
|
||||
def validate(body: CorrectWordFileBody, *, check_size: bool) -> None:
|
||||
if body.file_name is None or not body.file_name.strip():
|
||||
raise AppError(10034, "文件名不能为空")
|
||||
if not body.content:
|
||||
raise AppError(10034, "替换词内容不能为空")
|
||||
if check_size and body.file_size is not None and body.file_size > 1024 * 1024:
|
||||
raise AppError(10204)
|
||||
|
||||
async def create(self, body: CorrectWordFileBody, user: AuthUser) -> dict[str, Any]:
|
||||
self.validate(body, check_size=True)
|
||||
assert body.file_name is not None
|
||||
assert body.content is not None
|
||||
items = _parse_lines(body.content)
|
||||
file_id, now = uuid.uuid4().hex, shanghai_now_naive()
|
||||
values = {
|
||||
"id": file_id,
|
||||
"file_name": body.file_name,
|
||||
"word_count": len(items),
|
||||
"content": "\n".join(body.content),
|
||||
"creator": user.id,
|
||||
"now": now,
|
||||
}
|
||||
async with self.repository.session.begin():
|
||||
if await self.repository.name_exists(user.id, body.file_name):
|
||||
raise AppError(10203)
|
||||
await self.repository.insert_file(values)
|
||||
await self.repository.insert_items(
|
||||
[
|
||||
{"id": uuid.uuid4().hex, "file_id": file_id, "source_word": source, "target_word": target}
|
||||
for source, target in items
|
||||
]
|
||||
)
|
||||
return file_vo({**values, "created_at": now, "updated_at": None})
|
||||
|
||||
async def update(self, file_id: str, body: CorrectWordFileBody, user: AuthUser) -> None:
|
||||
self.validate(body, check_size=False)
|
||||
assert body.file_name is not None
|
||||
assert body.content is not None
|
||||
items = _parse_lines(body.content)
|
||||
async with self.repository.session.begin():
|
||||
row = await self.repository.get_file(file_id, for_update=True)
|
||||
if row is None:
|
||||
return
|
||||
if await self.repository.name_exists(user.id, body.file_name, file_id):
|
||||
raise AppError(500, f"文件名已存在:{body.file_name}")
|
||||
await self.repository.delete_items(file_id)
|
||||
await self.repository.insert_items(
|
||||
[
|
||||
{"id": uuid.uuid4().hex, "file_id": file_id, "source_word": source, "target_word": target}
|
||||
for source, target in items
|
||||
]
|
||||
)
|
||||
await self.repository.update_file(
|
||||
{
|
||||
"id": file_id,
|
||||
"file_name": body.file_name,
|
||||
"word_count": len(items),
|
||||
"content": "\n".join(body.content),
|
||||
"updater": user.id,
|
||||
"now": shanghai_now_naive(),
|
||||
}
|
||||
)
|
||||
|
||||
async def page(self, user: AuthUser, page: str | None, limit: str | None) -> dict[str, Any]:
|
||||
current, size = max(int(page or "1"), 1), int(limit or "10")
|
||||
rows, total = await self.repository.list_files(user.id, offset=(current - 1) * size, limit=size)
|
||||
return {"total": total, "list": [file_vo(row) for row in rows]}
|
||||
|
||||
async def all(self, user: AuthUser) -> list[dict[str, Any]]:
|
||||
rows, _ = await self.repository.list_files(user.id)
|
||||
return [file_vo(row) for row in rows]
|
||||
|
||||
async def get(self, file_id: str) -> dict[str, Any] | None:
|
||||
row = await self.repository.get_file(file_id)
|
||||
return file_vo(row) if row else None
|
||||
|
||||
async def delete(self, file_ids: list[str]) -> None:
|
||||
async with self.repository.session.begin():
|
||||
for file_id in file_ids:
|
||||
if file_id and file_id.strip():
|
||||
await self.repository.delete_file_graph(file_id.strip())
|
||||
@@ -0,0 +1,978 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import base64
|
||||
import hashlib
|
||||
import hmac
|
||||
import json
|
||||
import logging
|
||||
import random
|
||||
import re
|
||||
import secrets
|
||||
import uuid
|
||||
from collections.abc import Callable, Mapping, Sequence
|
||||
from datetime import datetime, timezone
|
||||
from pathlib import Path
|
||||
from typing import Any, cast
|
||||
from zoneinfo import ZoneInfo
|
||||
|
||||
import httpx
|
||||
from redis.asyncio import Redis
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
|
||||
from app.core.config import get_settings
|
||||
from app.core.database import get_session_factory
|
||||
from app.core.errors import AppError
|
||||
from app.core.redis import JavaRedisCodec, get_redis
|
||||
from app.core.security import AuthUser, shanghai_now_naive
|
||||
from app.integrations.mqtt_gateway import post_json
|
||||
from app.repositories.device import DeviceRepository
|
||||
from app.schemas.device import DeviceManualAddRequest, DeviceReportRequest, DeviceUpdateRequest, OtaRecord
|
||||
from app.services.system_params import SystemParamService
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
DEFAULT_TTL_SECONDS = 24 * 60 * 60
|
||||
INVALID_FIRMWARE_URL = (
|
||||
"http://xiaozhi.server.com:8002/xiaozhi/otaMag/download/NOT_ACTIVATED_FIRMWARE_THIS_IS_A_INVALID_URL"
|
||||
)
|
||||
MAC_PATTERN = re.compile(r"^([0-9A-Za-z]{2}[:-]){5}([0-9A-Za-z]{2})$")
|
||||
OTA_ORDER_COLUMNS = {
|
||||
"id": "id",
|
||||
"firmwareName": "firmware_name",
|
||||
"firmware_name": "firmware_name",
|
||||
"type": "type",
|
||||
"version": "version",
|
||||
"size": "size",
|
||||
"sort": "sort",
|
||||
"updateDate": "update_date",
|
||||
"update_date": "update_date",
|
||||
"createDate": "create_date",
|
||||
"create_date": "create_date",
|
||||
}
|
||||
|
||||
|
||||
def is_blank(value: str | None) -> bool:
|
||||
return value is None or not value.strip()
|
||||
|
||||
|
||||
def _java_semicolon_split(value: str) -> list[str]:
|
||||
parts = value.split(";")
|
||||
while parts and parts[-1] == "":
|
||||
parts.pop()
|
||||
return parts
|
||||
|
||||
|
||||
def _mapping(value: Any) -> dict[str, Any] | None:
|
||||
if isinstance(value, dict):
|
||||
if "@class" in value:
|
||||
return {str(key): item for key, item in value.items() if key != "@class"}
|
||||
return {str(key): item for key, item in value.items()}
|
||||
if isinstance(value, list) and len(value) == 2 and isinstance(value[1], dict):
|
||||
return {str(key): item for key, item in value[1].items()}
|
||||
return None
|
||||
|
||||
|
||||
async def redis_get(key: str, client: Redis | None = None) -> Any:
|
||||
selected = client or get_redis()
|
||||
raw = await cast(Any, selected.get(key))
|
||||
return JavaRedisCodec.decode(raw)
|
||||
|
||||
|
||||
async def redis_set(key: str, value: Any, *, ttl: int = DEFAULT_TTL_SECONDS, client: Redis | None = None) -> None:
|
||||
selected = client or get_redis()
|
||||
await cast(Any, selected.set(key, JavaRedisCodec.encode(value), ex=ttl))
|
||||
|
||||
|
||||
async def redis_delete(*keys: str, client: Redis | None = None) -> None:
|
||||
if not keys:
|
||||
return
|
||||
selected = client or get_redis()
|
||||
await cast(Any, selected.delete(*keys))
|
||||
|
||||
|
||||
async def redis_increment(key: str, *, ttl: int = DEFAULT_TTL_SECONDS, client: Redis | None = None) -> int:
|
||||
selected = client or get_redis()
|
||||
value = int(await cast(Any, selected.incr(key)))
|
||||
await cast(Any, selected.expire(key, ttl))
|
||||
return value
|
||||
|
||||
|
||||
class DeviceService:
|
||||
def __init__(
|
||||
self,
|
||||
session: AsyncSession,
|
||||
*,
|
||||
redis_client: Redis | None = None,
|
||||
http_client: httpx.AsyncClient | None = None,
|
||||
):
|
||||
self.session = session
|
||||
self.repository = DeviceRepository(session)
|
||||
self.params = SystemParamService(session)
|
||||
self.redis = redis_client
|
||||
self.http_client = http_client
|
||||
|
||||
async def register_device(self, mac_address: str) -> str:
|
||||
while True:
|
||||
code = f"{secrets.randbelow(1_000_000):06d}"
|
||||
key = f"sys:device:captcha:{code}"
|
||||
if is_blank(cast(str | None, await redis_get(key, self.redis))):
|
||||
await redis_set(key, mac_address, client=self.redis)
|
||||
return code
|
||||
|
||||
async def activate_bound_device(self, *, agent_id: str, activation_code: str, user: AuthUser) -> None:
|
||||
if is_blank(activation_code):
|
||||
raise AppError(10061)
|
||||
code_key = f"ota:activation:code:{activation_code}"
|
||||
device_id_value = await redis_get(code_key, self.redis)
|
||||
if device_id_value in (None, ""):
|
||||
raise AppError(10062)
|
||||
device_id = str(device_id_value)
|
||||
safe_device_id = device_id.replace(":", "_").lower()
|
||||
data_key = f"ota:activation:data:{safe_device_id}"
|
||||
cached = _mapping(await redis_get(data_key, self.redis))
|
||||
if cached is None or str(cached.get("activation_code") or "") != activation_code:
|
||||
raise AppError(10062)
|
||||
if await self.repository.get_device(device_id) is not None:
|
||||
raise AppError(10063)
|
||||
|
||||
now = shanghai_now_naive()
|
||||
values = {
|
||||
"id": device_id,
|
||||
"user_id": user.id,
|
||||
"mac_address": cached.get("mac_address"),
|
||||
"last_connected_at": now,
|
||||
"auto_update": 1,
|
||||
"board": cached.get("board"),
|
||||
"alias": None,
|
||||
"agent_id": agent_id,
|
||||
"app_version": cached.get("app_version"),
|
||||
"sort": None,
|
||||
"updater": user.id,
|
||||
"update_date": now,
|
||||
"creator": user.id,
|
||||
"create_date": now,
|
||||
}
|
||||
try:
|
||||
await self.repository.insert_device(values)
|
||||
await self.session.commit()
|
||||
except Exception:
|
||||
await self.session.rollback()
|
||||
raise
|
||||
await redis_delete(data_key, code_key, f"agent:device:count:{agent_id}", client=self.redis)
|
||||
|
||||
async def list_user_devices(self, user_id: int, agent_id: str) -> list[dict[str, Any]]:
|
||||
devices = await self.repository.get_user_devices(user_id, agent_id)
|
||||
return [self._user_device_view(row) for row in devices]
|
||||
|
||||
async def get_online_data(self, agent_id: str, user: AuthUser) -> str:
|
||||
gateway = await self.params.get_value("server.mqtt_manager_api", from_cache=True)
|
||||
if is_blank(gateway) or gateway == "null":
|
||||
return ""
|
||||
devices = await self.repository.get_user_devices(user.id, agent_id)
|
||||
client_ids = {
|
||||
self._mqtt_client_id(
|
||||
str(device.get("board") or "GID_default"),
|
||||
str(device.get("mac_address") or "unknown"),
|
||||
)
|
||||
for device in devices
|
||||
}
|
||||
if not client_ids:
|
||||
return ""
|
||||
signature_key = await self.params.get_value("server.mqtt_signature_key", from_cache=False)
|
||||
return await post_json(
|
||||
f"http://{gateway}/api/devices/status",
|
||||
{"clientIds": sorted(client_ids)},
|
||||
signature_key or "",
|
||||
timeout_seconds=get_settings().external_request_timeout_seconds,
|
||||
client=self.http_client,
|
||||
)
|
||||
|
||||
async def unbind(self, *, user_id: int, device_id: str) -> None:
|
||||
device = await self.repository.get_device(device_id)
|
||||
if device is None:
|
||||
return
|
||||
mac_address = device.get("mac_address")
|
||||
agent_id = device.get("agent_id")
|
||||
if not is_blank(None if agent_id is None else str(agent_id)):
|
||||
await redis_delete(f"agent:device:count:{agent_id}", client=self.redis)
|
||||
try:
|
||||
await self.repository.delete_device_for_user(device_id, user_id)
|
||||
await self.session.commit()
|
||||
except Exception:
|
||||
await self.session.rollback()
|
||||
raise
|
||||
try:
|
||||
if mac_address is not None:
|
||||
await self.repository.delete_address_books_for_macs([str(mac_address)])
|
||||
await self.session.commit()
|
||||
except Exception:
|
||||
await self.session.rollback()
|
||||
raise
|
||||
await self.refresh_address_book_cache()
|
||||
|
||||
async def update_device(
|
||||
self,
|
||||
*,
|
||||
device_id: str,
|
||||
request: DeviceUpdateRequest,
|
||||
user: AuthUser,
|
||||
) -> bool:
|
||||
device = await self.repository.get_device(device_id)
|
||||
if device is None or int(device.get("user_id") or -1) != user.id:
|
||||
return False
|
||||
await self.repository.update_device_info(
|
||||
device_id,
|
||||
auto_update=request.auto_update,
|
||||
alias=request.alias,
|
||||
updater=user.id,
|
||||
now=shanghai_now_naive(),
|
||||
)
|
||||
await self.session.commit()
|
||||
return True
|
||||
|
||||
async def manual_add(self, *, request: DeviceManualAddRequest, user: AuthUser) -> None:
|
||||
mac_address = request.mac_address
|
||||
if mac_address is not None and await self.repository.get_device_by_mac(mac_address) is not None:
|
||||
raise AppError(10161)
|
||||
now = shanghai_now_naive()
|
||||
values = {
|
||||
"id": uuid.uuid4().hex if mac_address in (None, "") else mac_address,
|
||||
"user_id": user.id,
|
||||
"mac_address": mac_address,
|
||||
"last_connected_at": now,
|
||||
"auto_update": 1,
|
||||
"board": request.board,
|
||||
"alias": None,
|
||||
"agent_id": request.agent_id,
|
||||
"app_version": request.app_version,
|
||||
"sort": None,
|
||||
"updater": user.id,
|
||||
"update_date": now,
|
||||
"creator": user.id,
|
||||
"create_date": now,
|
||||
}
|
||||
try:
|
||||
await self.repository.insert_device(values)
|
||||
await self.session.commit()
|
||||
except Exception:
|
||||
await self.session.rollback()
|
||||
raise
|
||||
agent_cache_id = "null" if request.agent_id is None else request.agent_id
|
||||
await redis_delete(f"agent:device:count:{agent_cache_id}", client=self.redis)
|
||||
|
||||
async def get_tools(self, *, device_id: str, user: AuthUser) -> dict[str, Any] | None:
|
||||
gateway_and_device = await self._gateway_device(device_id, user)
|
||||
if gateway_and_device is None:
|
||||
return None
|
||||
gateway, device = gateway_and_device
|
||||
client_id = self._mqtt_client_id(
|
||||
str(device.get("board") or "GID_default"),
|
||||
str(device.get("mac_address") or "unknown"),
|
||||
)
|
||||
url = f"http://{gateway}/api/commands/{client_id}"
|
||||
all_tools: list[Any] = []
|
||||
cursor: str | None = None
|
||||
while True:
|
||||
params: dict[str, Any] = {"withUserTools": True}
|
||||
if cursor is not None and cursor.strip():
|
||||
params["cursor"] = cursor
|
||||
body = {
|
||||
"type": "mcp",
|
||||
"payload": {"jsonrpc": "2.0", "id": 2, "method": "tools/list", "params": params},
|
||||
}
|
||||
response_body = await self._post_gateway(url, body)
|
||||
if is_blank(response_body):
|
||||
break
|
||||
payload = json.loads(response_body)
|
||||
if not isinstance(payload, dict) or not bool(payload.get("success", False)):
|
||||
break
|
||||
data = payload.get("data")
|
||||
if not isinstance(data, dict):
|
||||
break
|
||||
tools = data.get("tools")
|
||||
if isinstance(tools, list):
|
||||
all_tools.extend(tools)
|
||||
next_cursor = data.get("nextCursor")
|
||||
if not isinstance(next_cursor, str) or not next_cursor.strip():
|
||||
break
|
||||
cursor = next_cursor
|
||||
return None if not all_tools else {"tools": all_tools}
|
||||
|
||||
async def call_tool(
|
||||
self,
|
||||
*,
|
||||
device_id: str,
|
||||
tool_name: str,
|
||||
arguments: dict[str, Any] | None,
|
||||
user: AuthUser,
|
||||
) -> Any:
|
||||
gateway_and_device = await self._gateway_device(device_id, user)
|
||||
if gateway_and_device is None:
|
||||
return None
|
||||
gateway, device = gateway_and_device
|
||||
client_id = self._mqtt_client_id(
|
||||
str(device.get("board") or "GID_default"),
|
||||
str(device.get("mac_address") or "unknown"),
|
||||
)
|
||||
response_body = await self._post_gateway(
|
||||
f"http://{gateway}/api/commands/{client_id}",
|
||||
{
|
||||
"type": "mcp",
|
||||
"payload": {
|
||||
"jsonrpc": "2.0",
|
||||
"id": 2,
|
||||
"method": "tools/call",
|
||||
"params": {"name": tool_name, "arguments": arguments},
|
||||
},
|
||||
},
|
||||
)
|
||||
if is_blank(response_body):
|
||||
return None
|
||||
payload = json.loads(response_body)
|
||||
if not isinstance(payload, dict) or not bool(payload.get("success", False)):
|
||||
return None
|
||||
data = payload.get("data")
|
||||
content = data.get("content") if isinstance(data, dict) else None
|
||||
if not isinstance(content, list) or not content or not isinstance(content[0], dict):
|
||||
return None
|
||||
first = content[0]
|
||||
if first.get("type") != "text" or not isinstance(first.get("text"), str):
|
||||
return None
|
||||
text = str(first["text"])
|
||||
if not text.strip():
|
||||
return None
|
||||
trimmed = text.strip()
|
||||
if trimmed.startswith("{") or trimmed.startswith("["):
|
||||
try:
|
||||
parsed = json.loads(trimmed)
|
||||
return parsed if isinstance(parsed, dict) else trimmed
|
||||
except json.JSONDecodeError:
|
||||
return trimmed
|
||||
if trimmed == "true":
|
||||
return True
|
||||
if trimmed == "false":
|
||||
return False
|
||||
return trimmed
|
||||
|
||||
async def address_book(self, mac_address: str) -> list[dict[str, Any]]:
|
||||
rows = await self.repository.get_address_book(mac_address)
|
||||
for row in rows:
|
||||
if row.get("has_permission") is not None:
|
||||
row["has_permission"] = bool(row["has_permission"])
|
||||
return rows
|
||||
|
||||
async def lookup_address_book(self, *, caller_mac: str, nickname: str) -> dict[str, str | None] | None:
|
||||
books = await self.all_address_books()
|
||||
caller_book = books.get(caller_mac.lower())
|
||||
if caller_book is None:
|
||||
return None
|
||||
target_with_permission = caller_book.get(nickname)
|
||||
if target_with_permission is None:
|
||||
return None
|
||||
parts = target_with_permission.split("|")
|
||||
target_mac = parts[0]
|
||||
has_permission = len(parts) > 1 and parts[1] == "1"
|
||||
target_book = books.get(target_mac.lower())
|
||||
if target_book is None:
|
||||
return None
|
||||
caller_nickname = target_book.get(caller_mac.lower())
|
||||
return {
|
||||
"targetMac": target_mac,
|
||||
"callerNickname": caller_nickname,
|
||||
"hasPermission": "true" if has_permission else "false",
|
||||
}
|
||||
|
||||
async def call_by_nickname(self, *, caller_mac: str, nickname: str, answer: bool) -> dict[str, Any]:
|
||||
books = await self.all_address_books()
|
||||
if answer:
|
||||
return await self._post_call("/api/call/accept", {"mac": caller_mac}, "接听")
|
||||
caller_book = books.get(caller_mac.lower())
|
||||
if caller_book is None or nickname not in caller_book:
|
||||
return {"status": "error", "message": f"未找到备注为'{nickname}'的设备"}
|
||||
parts = caller_book[nickname].split("|")
|
||||
target_mac = parts[0]
|
||||
allowed = len(parts) > 1 and parts[1] == "1"
|
||||
if not allowed:
|
||||
return {"status": "error", "message": "呼叫失败,您没有权限呼叫该设备"}
|
||||
target_book = books.get(target_mac.lower())
|
||||
caller_nickname = target_book.get(caller_mac.lower()) if target_book is not None else None
|
||||
if is_blank(caller_nickname):
|
||||
caller = await self.repository.get_device_by_mac(caller_mac)
|
||||
if caller is None:
|
||||
raise RuntimeError("caller device does not exist")
|
||||
caller_nickname = None if caller.get("alias") is None else str(caller["alias"])
|
||||
if is_blank(caller_nickname):
|
||||
caller_nickname = self._mac_device_name(caller_mac)
|
||||
return await self._post_call(
|
||||
"/api/call/request",
|
||||
{"caller_mac": caller_mac, "target_mac": target_mac, "caller_nickname": caller_nickname},
|
||||
"呼叫",
|
||||
)
|
||||
|
||||
async def save_address_book(
|
||||
self,
|
||||
*,
|
||||
mac_address: str,
|
||||
target_mac: str,
|
||||
alias: str | None,
|
||||
has_permission: bool | None,
|
||||
actor: int,
|
||||
) -> None:
|
||||
record = await self.repository.get_address_book_record(mac_address, target_mac)
|
||||
now = shanghai_now_naive()
|
||||
if record is None:
|
||||
final_alias = alias
|
||||
if is_blank(final_alias):
|
||||
target = await self.repository.get_device_by_mac(target_mac)
|
||||
if target is None:
|
||||
raise RuntimeError("target device does not exist")
|
||||
final_alias = None if target.get("alias") is None else str(target["alias"])
|
||||
final_alias = await self._unique_alias(mac_address, final_alias)
|
||||
await self.repository.insert_address_book(
|
||||
mac_address=mac_address,
|
||||
target_mac=target_mac,
|
||||
alias=final_alias,
|
||||
has_permission=has_permission,
|
||||
actor=actor,
|
||||
now=now,
|
||||
)
|
||||
await self.session.commit()
|
||||
else:
|
||||
if alias is not None:
|
||||
await self.repository.update_address_alias(
|
||||
mac_address,
|
||||
target_mac,
|
||||
await self._unique_alias(mac_address, alias),
|
||||
now=now,
|
||||
)
|
||||
await self.session.commit()
|
||||
await self.refresh_address_book_cache()
|
||||
if has_permission is not None:
|
||||
await self.repository.update_address_permission(
|
||||
mac_address,
|
||||
target_mac,
|
||||
has_permission,
|
||||
now=now,
|
||||
)
|
||||
await self.session.commit()
|
||||
await self.refresh_address_book_cache()
|
||||
|
||||
async def all_address_books(self) -> dict[str, dict[str, str]]:
|
||||
cached = _mapping(await redis_get("device:address_book:all", self.redis))
|
||||
if cached is not None:
|
||||
result: dict[str, dict[str, str]] = {}
|
||||
for key, value in cached.items():
|
||||
nested = _mapping(value)
|
||||
if nested is not None:
|
||||
result[key] = {str(field): str(item) for field, item in nested.items()}
|
||||
return result
|
||||
return await self.refresh_address_book_cache()
|
||||
|
||||
async def refresh_address_book_cache(self) -> dict[str, dict[str, str]]:
|
||||
records = await self.repository.get_all_address_book()
|
||||
result: dict[str, dict[str, str]] = {}
|
||||
reverse: dict[str, str] = {}
|
||||
for record in records:
|
||||
mac_a = str(record["mac_address"]).lower()
|
||||
mac_b = str(record["target_mac"]).lower()
|
||||
alias = record.get("alias")
|
||||
if alias not in (None, ""):
|
||||
alias_string = str(alias)
|
||||
result.setdefault(mac_a, {})[alias_string] = (
|
||||
f"{mac_b}|{'1' if bool(record.get('has_permission')) else '0'}"
|
||||
)
|
||||
reverse[f"{mac_b}:{mac_a}"] = alias_string
|
||||
for record in records:
|
||||
mac_a = str(record["mac_address"]).lower()
|
||||
mac_b = str(record["target_mac"]).lower()
|
||||
reverse_alias = reverse.get(f"{mac_a}:{mac_b}")
|
||||
if isinstance(reverse_alias, str) and reverse_alias:
|
||||
result.setdefault(mac_b, {})[mac_a] = reverse_alias
|
||||
await redis_set("device:address_book:all", result, client=self.redis)
|
||||
return result
|
||||
|
||||
async def check_ota(
|
||||
self,
|
||||
*,
|
||||
device_id: str,
|
||||
client_id: str,
|
||||
report: DeviceReportRequest,
|
||||
request_url: str,
|
||||
client_ip: str,
|
||||
defer_connection_update: Callable[[str, str | None, str | None], None] | None = None,
|
||||
) -> dict[str, Any]:
|
||||
now = datetime.now(tz=ZoneInfo(get_settings().timezone))
|
||||
utc_offset = now.utcoffset()
|
||||
response: dict[str, Any] = {
|
||||
"server_time": {
|
||||
"timestamp": int(now.timestamp() * 1000),
|
||||
"timeZone": get_settings().timezone,
|
||||
"timezone_offset": int((utc_offset.total_seconds() if utc_offset is not None else 0) / 60),
|
||||
},
|
||||
"activation": None,
|
||||
"error": None,
|
||||
"firmware": None,
|
||||
"websocket": None,
|
||||
"mqtt": None,
|
||||
}
|
||||
device = await self.repository.get_device_by_mac(device_id)
|
||||
if device is None:
|
||||
if report.application is None:
|
||||
raise RuntimeError("application is required")
|
||||
response["firmware"] = {
|
||||
"version": report.application.version,
|
||||
"url": INVALID_FIRMWARE_URL,
|
||||
}
|
||||
elif device.get("auto_update") is None:
|
||||
raise RuntimeError("auto_update is null")
|
||||
elif int(device["auto_update"]) != 0:
|
||||
ota_type = report.board.type if report.board is not None else None
|
||||
current_version = report.application.version if report.application is not None else None
|
||||
response["firmware"] = await self._firmware_info(ota_type, current_version, request_url)
|
||||
|
||||
websocket_url = await self.params.get_value("server.websocket", from_cache=True)
|
||||
auth_enabled = await self.params.get_value("server.auth.enabled", from_cache=True)
|
||||
websocket_token = ""
|
||||
if (auth_enabled or "").lower() == "true":
|
||||
try:
|
||||
websocket_token = await self._websocket_token(client_id, device_id)
|
||||
except Exception:
|
||||
logger.exception("WebSocket token generation failed")
|
||||
if is_blank(websocket_url) or websocket_url == "null":
|
||||
selected_websocket = "ws://xiaozhi.server.com:8000/xiaozhi/v1/"
|
||||
else:
|
||||
websocket_urls = _java_semicolon_split(websocket_url or "")
|
||||
selected_websocket = (
|
||||
random.choice(websocket_urls) # noqa: S311
|
||||
if websocket_urls
|
||||
else "ws://xiaozhi.server.com:8000/xiaozhi/v1/"
|
||||
)
|
||||
response["websocket"] = {"url": selected_websocket, "token": websocket_token}
|
||||
|
||||
mqtt_endpoint = await self.params.get_value("server.mqtt_gateway", from_cache=True)
|
||||
if mqtt_endpoint not in (None, "", "null"):
|
||||
try:
|
||||
group_id = str(device.get("board") or "GID_default") if device is not None else "GID_default"
|
||||
mqtt = await self._mqtt_config(device_id, group_id, client_ip)
|
||||
if mqtt is not None:
|
||||
mqtt["endpoint"] = mqtt_endpoint
|
||||
response["mqtt"] = mqtt
|
||||
except Exception:
|
||||
logger.exception("MQTT credential generation failed")
|
||||
|
||||
if device is None:
|
||||
response["activation"] = await self._activation(device_id, report)
|
||||
else:
|
||||
app_version = report.application.version if report.application is not None else None
|
||||
agent_id = device.get("agent_id")
|
||||
normalized_agent_id = None if agent_id is None else str(agent_id)
|
||||
if defer_connection_update is not None:
|
||||
defer_connection_update(str(device["id"]), normalized_agent_id, app_version)
|
||||
else:
|
||||
try:
|
||||
await self._persist_connection_update(
|
||||
str(device["id"]),
|
||||
normalized_agent_id,
|
||||
app_version,
|
||||
)
|
||||
except Exception:
|
||||
logger.exception("Asynchronous device connection update failed")
|
||||
return cast(dict[str, Any], self._drop_none(response))
|
||||
|
||||
async def _persist_connection_update(
|
||||
self,
|
||||
device_id: str,
|
||||
agent_id: str | None,
|
||||
app_version: str | None,
|
||||
) -> None:
|
||||
connection_time = shanghai_now_naive()
|
||||
try:
|
||||
await self.repository.update_connection(device_id, app_version=app_version, now=connection_time)
|
||||
await self.session.commit()
|
||||
except Exception:
|
||||
await self.session.rollback()
|
||||
raise
|
||||
if not is_blank(agent_id):
|
||||
await redis_set(f"agent:device:lastConnected:{agent_id}", connection_time, client=self.redis)
|
||||
|
||||
@staticmethod
|
||||
async def persist_connection_update_background(
|
||||
device_id: str,
|
||||
agent_id: str | None,
|
||||
app_version: str | None,
|
||||
) -> None:
|
||||
try:
|
||||
async with get_session_factory()() as session:
|
||||
await DeviceService(session)._persist_connection_update(device_id, agent_id, app_version)
|
||||
except Exception:
|
||||
logger.exception("Asynchronous device connection update failed")
|
||||
|
||||
async def ota_health_text(self) -> str:
|
||||
mqtt_gateway = await self.params.get_value("server.mqtt_gateway", from_cache=False)
|
||||
if is_blank(mqtt_gateway):
|
||||
return "OTA接口不正常,缺少mqtt_gateway地址,请登录智控台,在参数管理找到【server.mqtt_gateway】配置"
|
||||
websocket = await self.params.get_value("server.websocket", from_cache=True)
|
||||
if is_blank(websocket) or websocket == "null":
|
||||
return "OTA接口不正常,缺少websocket地址,请登录智控台,在参数管理找到【server.websocket】配置"
|
||||
ota_url = await self.params.get_value("server.ota", from_cache=True)
|
||||
if is_blank(ota_url) or ota_url == "null":
|
||||
return "OTA接口不正常,缺少ota地址,请登录智控台,在参数管理找到【server.ota】配置"
|
||||
return f"OTA接口运行正常,websocket集群数量:{len(_java_semicolon_split(websocket or ''))}"
|
||||
|
||||
async def ota_page(self, query: Mapping[str, Any]) -> dict[str, Any]:
|
||||
page = self._positive_int(query.get("page"), 1)
|
||||
limit = self._positive_int(query.get("limit"), 10)
|
||||
requested = query.get("orderField")
|
||||
requested_fields = [requested] if isinstance(requested, str) else list(requested or [])
|
||||
fields = [OTA_ORDER_COLUMNS[field] for field in requested_fields if field in OTA_ORDER_COLUMNS]
|
||||
if not fields:
|
||||
fields = ["update_date"]
|
||||
ascending = str(query.get("order") or "").lower() == "asc" if requested_fields else True
|
||||
firmware_name = query.get("firmwareName")
|
||||
name = str(firmware_name) if firmware_name is not None else None
|
||||
rows = await self.repository.list_ota(
|
||||
page=page,
|
||||
limit=limit,
|
||||
firmware_name=name,
|
||||
order_fields=fields,
|
||||
ascending=ascending,
|
||||
)
|
||||
rows = [self._ota_response_record(row) for row in rows]
|
||||
return {"total": await self.repository.count_ota(name), "list": rows}
|
||||
|
||||
async def get_ota_record(self, ota_id: str) -> dict[str, Any] | None:
|
||||
row = await self.repository.get_ota(ota_id)
|
||||
return None if row is None else self._ota_response_record(row)
|
||||
|
||||
async def save_ota(self, record: OtaRecord, user: AuthUser) -> None:
|
||||
values = record.model_dump(by_alias=False)
|
||||
existing = await self.repository.get_first_ota_by_type(record.type or "")
|
||||
now = shanghai_now_naive()
|
||||
if existing is not None:
|
||||
values["updater"] = record.updater if record.updater is not None else user.id
|
||||
values["update_date"] = record.update_date if record.update_date is not None else now
|
||||
await self.repository.update_ota(str(existing["id"]), values)
|
||||
else:
|
||||
values["id"] = record.id or uuid.uuid4().hex
|
||||
values["creator"] = record.creator if record.creator is not None else user.id
|
||||
values["updater"] = record.updater if record.updater is not None else user.id
|
||||
values["create_date"] = record.create_date if record.create_date is not None else now
|
||||
values["update_date"] = record.update_date if record.update_date is not None else now
|
||||
await self.repository.insert_ota(values)
|
||||
await self.session.commit()
|
||||
|
||||
async def update_ota(self, ota_id: str, record: OtaRecord, user: AuthUser) -> None:
|
||||
if await self.repository.count_duplicate_ota(
|
||||
ota_id=ota_id,
|
||||
ota_type=record.type,
|
||||
version=record.version,
|
||||
):
|
||||
raise RuntimeError("已存在相同类型和版本的固件,请修改后重试")
|
||||
values = record.model_dump(by_alias=False)
|
||||
values["updater"] = record.updater if record.updater is not None else user.id
|
||||
values["update_date"] = shanghai_now_naive()
|
||||
await self.repository.update_ota(ota_id, values)
|
||||
await self.session.commit()
|
||||
|
||||
async def delete_ota(self, ids: Sequence[str]) -> None:
|
||||
await self.repository.delete_ota(ids)
|
||||
await self.session.commit()
|
||||
|
||||
async def create_ota_download_id(self, ota_id: str) -> str:
|
||||
value = str(uuid.uuid4())
|
||||
await redis_set(f"ota:id:{value}", ota_id, client=self.redis)
|
||||
return value
|
||||
|
||||
async def resolve_ota_download(self, download_id: str) -> tuple[Path, str] | None:
|
||||
id_key = f"ota:id:{download_id}"
|
||||
ota_value = await redis_get(id_key, self.redis)
|
||||
if is_blank(None if ota_value is None else str(ota_value)):
|
||||
return None
|
||||
count_key = f"ota:download:count:{download_id}"
|
||||
count_value = await redis_get(count_key, self.redis)
|
||||
count = int(count_value or 0)
|
||||
if count >= 3:
|
||||
await redis_delete(count_key, id_key, client=self.redis)
|
||||
return None
|
||||
await redis_set(count_key, count + 1, client=self.redis)
|
||||
|
||||
ota_id = str(ota_value)
|
||||
if ota_id.startswith("file:"):
|
||||
firmware_path = ota_id[5:]
|
||||
ota_type = "assets"
|
||||
version = "1.0.0"
|
||||
else:
|
||||
record = await self.repository.get_ota(ota_id)
|
||||
firmware_value = None if record is None else record.get("firmware_path")
|
||||
if record is None or is_blank(None if firmware_value is None else str(firmware_value)):
|
||||
return None
|
||||
firmware_path = str(record["firmware_path"])
|
||||
ota_type = str(record.get("type"))
|
||||
version = str(record.get("version"))
|
||||
raw_path = Path(firmware_path)
|
||||
candidates = [raw_path] if raw_path.is_absolute() else [Path.cwd() / raw_path]
|
||||
if not raw_path.is_absolute() and raw_path.parts and raw_path.parts[0] == "uploadfile":
|
||||
candidates.insert(0, get_settings().upload_dir.joinpath(*raw_path.parts[1:]))
|
||||
candidates.append(Path.cwd() / "firmware" / raw_path.name)
|
||||
resolved = next((candidate for candidate in candidates if candidate.is_file()), None)
|
||||
if resolved is None:
|
||||
return None
|
||||
original_name = f"{ota_type}_{version}"
|
||||
dot_index = firmware_path.rfind(".")
|
||||
if dot_index >= 0:
|
||||
original_name += firmware_path[dot_index:]
|
||||
safe_name = re.sub(r"[^a-zA-Z0-9._-]", "_", original_name)
|
||||
return resolved, safe_name
|
||||
|
||||
async def save_firmware_file(self, *, filename: str | None, content: bytes) -> str:
|
||||
if not content:
|
||||
raise ValueError("上传文件不能为空")
|
||||
if filename is None:
|
||||
raise ValueError("文件名不能为空")
|
||||
dot_index = filename.rfind(".")
|
||||
if dot_index < 0:
|
||||
raise RuntimeError("文件名缺少扩展名")
|
||||
extension = filename[dot_index:].lower()
|
||||
if extension not in {".bin", ".apk"}:
|
||||
raise ValueError("只允许上传.bin和.apk格式的文件")
|
||||
digest = hashlib.md5(content, usedforsecurity=False).hexdigest()
|
||||
directory = get_settings().upload_dir
|
||||
directory.mkdir(parents=True, exist_ok=True)
|
||||
filename_on_disk = f"{digest}{extension}"
|
||||
physical_path = directory / filename_on_disk
|
||||
if not physical_path.exists():
|
||||
with physical_path.open("xb") as stream:
|
||||
stream.write(content)
|
||||
# Keep Java's database/API value stable even when the physical upload
|
||||
# volume is mounted elsewhere (for example /data/uploads in Docker).
|
||||
return str(Path("uploadfile") / filename_on_disk)
|
||||
|
||||
async def save_assets_file(self, *, filename: str | None, content: bytes, user: AuthUser) -> str:
|
||||
ota_url = await self.params.get_value("server.ota", from_cache=True)
|
||||
if is_blank(ota_url) or ota_url == "null":
|
||||
raise AppError(10102)
|
||||
if len(content) > 20 * 1024 * 1024:
|
||||
raise AppError(10142)
|
||||
if not user.is_super_admin:
|
||||
count_key = f"ota:upload:count:{user.id}"
|
||||
current = int(await redis_get(count_key, self.redis) or 0)
|
||||
if current >= 50:
|
||||
raise AppError(10195)
|
||||
await redis_increment(count_key, client=self.redis)
|
||||
path = await self.save_firmware_file(filename=filename, content=content)
|
||||
download_id = await self.create_ota_download_id(f"file:{path}")
|
||||
return (ota_url or "").replace("/ota/", "/otaMag/download/") + download_id
|
||||
|
||||
async def _gateway_device(self, device_id: str, user: AuthUser) -> tuple[str, dict[str, Any]] | None:
|
||||
gateway = await self.params.get_value("server.mqtt_manager_api", from_cache=True)
|
||||
if is_blank(gateway) or gateway == "null":
|
||||
return None
|
||||
device = await self.repository.get_device(device_id)
|
||||
if device is None or int(device.get("user_id") or -1) != user.id:
|
||||
return None
|
||||
return gateway or "", device
|
||||
|
||||
async def _post_gateway(self, url: str, body: Any, *, timeout_seconds: float | None = None) -> str:
|
||||
key = await self.params.get_value("server.mqtt_signature_key", from_cache=False)
|
||||
return await post_json(
|
||||
url,
|
||||
body,
|
||||
key or "",
|
||||
timeout_seconds=timeout_seconds or get_settings().external_request_timeout_seconds,
|
||||
client=self.http_client,
|
||||
)
|
||||
|
||||
async def _post_call(self, path: str, body: dict[str, Any], action: str) -> dict[str, Any]:
|
||||
gateway = await self.params.get_value("server.mqtt_manager_api", from_cache=True)
|
||||
key = await self.params.get_value("server.mqtt_signature_key", from_cache=True)
|
||||
if is_blank(gateway) or gateway == "null" or is_blank(key) or (key or "").strip().lower() == "null":
|
||||
return {"status": "error", "message": f"{action}失败,网关配置缺失"}
|
||||
result: dict[str, Any] = {"status": "error"}
|
||||
try:
|
||||
text = await post_json(
|
||||
f"http://{gateway}{path}",
|
||||
body,
|
||||
key or "",
|
||||
timeout_seconds=5.0,
|
||||
client=self.http_client,
|
||||
)
|
||||
if text.strip():
|
||||
payload = json.loads(text)
|
||||
if isinstance(payload, dict):
|
||||
result["status"] = payload.get("status")
|
||||
result["message"] = payload.get("message")
|
||||
return result
|
||||
except Exception:
|
||||
return {"status": "error", "message": f"{action}失败,请稍后再试"}
|
||||
|
||||
async def _firmware_info(
|
||||
self,
|
||||
ota_type: str | None,
|
||||
current_version: str | None,
|
||||
request_url: str,
|
||||
) -> dict[str, Any] | None:
|
||||
if is_blank(ota_type):
|
||||
return None
|
||||
selected_version = current_version if not is_blank(current_version) else "0.0.0"
|
||||
ota = await self.repository.get_latest_ota(ota_type or "")
|
||||
download_url: str | None = None
|
||||
if ota is not None and self._compare_versions(ota.get("version"), selected_version) > 0:
|
||||
ota_url = await self.params.get_value("server.ota", from_cache=True)
|
||||
if is_blank(ota_url) or ota_url == "null":
|
||||
ota_url = request_url
|
||||
download_id = await self.create_ota_download_id(str(ota["id"]))
|
||||
download_url = (ota_url or "").replace("/ota/", "/otaMag/download/") + download_id
|
||||
return {
|
||||
"version": selected_version if ota is None else ota.get("version"),
|
||||
"url": download_url or INVALID_FIRMWARE_URL,
|
||||
}
|
||||
|
||||
async def _activation(self, device_id: str, report: DeviceReportRequest) -> dict[str, Any]:
|
||||
safe_device_id = device_id.replace(":", "_").lower()
|
||||
data_key = f"ota:activation:data:{safe_device_id}"
|
||||
cached = _mapping(await redis_get(data_key, self.redis))
|
||||
code = str(cached.get("activation_code")) if cached and cached.get("activation_code") is not None else None
|
||||
frontend = await self.params.get_value("server.fronted_url", from_cache=True)
|
||||
if code is None or not code.strip():
|
||||
code = f"{secrets.randbelow(1_000_000):06d}"
|
||||
board = (
|
||||
report.board.type
|
||||
if report.board is not None and report.board.type is not None
|
||||
else (report.chip_model_name or "unknown")
|
||||
)
|
||||
app_version = report.application.version if report.application is not None else None
|
||||
await redis_set(
|
||||
data_key,
|
||||
{
|
||||
"id": device_id,
|
||||
"mac_address": device_id,
|
||||
"board": board,
|
||||
"app_version": app_version,
|
||||
"deviceId": device_id,
|
||||
"activation_code": code,
|
||||
},
|
||||
client=self.redis,
|
||||
)
|
||||
await redis_set(f"ota:activation:code:{code}", device_id, client=self.redis)
|
||||
return {
|
||||
"code": code,
|
||||
"message": f"{frontend if frontend is not None else 'null'}\n{code}",
|
||||
"challenge": device_id,
|
||||
}
|
||||
|
||||
async def _websocket_token(self, client_id: str, username: str) -> str:
|
||||
secret = await self.params.get_value("server.secret", from_cache=False)
|
||||
if is_blank(secret):
|
||||
raise RuntimeError("WebSocket认证密钥未配置(server.secret)")
|
||||
timestamp = int(datetime.now().timestamp())
|
||||
message = f"{client_id}|{username}|{timestamp}".encode()
|
||||
signature = hmac.new((secret or "").encode(), message, hashlib.sha256).digest()
|
||||
encoded = base64.urlsafe_b64encode(signature).decode().rstrip("=")
|
||||
return f"{encoded}.{timestamp}"
|
||||
|
||||
async def _mqtt_config(self, mac_address: str, group_id: str, client_ip: str) -> dict[str, Any] | None:
|
||||
key = await self.params.get_value("server.mqtt_signature_key", from_cache=True)
|
||||
if is_blank(key):
|
||||
return None
|
||||
client_id = self._mqtt_client_id(group_id, mac_address)
|
||||
user_data = json.dumps({"ip": client_ip}, ensure_ascii=False, separators=(",", ":"))
|
||||
username = base64.b64encode(user_data.encode()).decode()
|
||||
password = base64.b64encode(
|
||||
hmac.new((key or "").encode(), f"{client_id}|{username}".encode(), hashlib.sha256).digest()
|
||||
).decode()
|
||||
safe_mac = mac_address.replace(":", "_")
|
||||
return {
|
||||
"client_id": client_id,
|
||||
"username": username,
|
||||
"password": password,
|
||||
"publish_topic": "device-server",
|
||||
"subscribe_topic": f"devices/p2p/{safe_mac}",
|
||||
}
|
||||
|
||||
@staticmethod
|
||||
def _mqtt_client_id(group_id: str, mac_address: str) -> str:
|
||||
safe_group = group_id.replace(":", "_")
|
||||
safe_mac = mac_address.replace(":", "_")
|
||||
return f"{safe_group}@@@{safe_mac}@@@{safe_mac}"
|
||||
|
||||
@staticmethod
|
||||
def _compare_versions(first: Any, second: Any) -> int:
|
||||
if first is None or second is None:
|
||||
return 0
|
||||
first = str(first)
|
||||
second = str(second)
|
||||
first_parts = first.split(".")
|
||||
second_parts = second.split(".")
|
||||
for index in range(max(len(first_parts), len(second_parts))):
|
||||
first_value = int(first_parts[index]) if index < len(first_parts) else 0
|
||||
second_value = int(second_parts[index]) if index < len(second_parts) else 0
|
||||
if first_value != second_value:
|
||||
return 1 if first_value > second_value else -1
|
||||
return 0
|
||||
|
||||
async def _unique_alias(self, mac_address: str, alias: str | None) -> str | None:
|
||||
existing = await self.repository.get_aliases(mac_address)
|
||||
if alias not in existing:
|
||||
return alias
|
||||
suffix = 1
|
||||
while f"{alias}{suffix}" in existing:
|
||||
suffix += 1
|
||||
return f"{alias}{suffix}"
|
||||
|
||||
@staticmethod
|
||||
def _mac_device_name(mac: str) -> str:
|
||||
return mac if len(mac) < 2 else f"尾号为{mac[-2:]}的设备"
|
||||
|
||||
@staticmethod
|
||||
def _positive_int(value: Any, default: int) -> int:
|
||||
if value is None:
|
||||
return default
|
||||
return int(str(value))
|
||||
|
||||
@staticmethod
|
||||
def _drop_none(value: Any) -> Any:
|
||||
if isinstance(value, dict):
|
||||
return {key: DeviceService._drop_none(item) for key, item in value.items() if item is not None}
|
||||
if isinstance(value, list):
|
||||
return [DeviceService._drop_none(item) for item in value]
|
||||
return value
|
||||
|
||||
@staticmethod
|
||||
def _user_device_view(row: Mapping[str, Any]) -> dict[str, Any]:
|
||||
return {
|
||||
"app_version": row.get("app_version"),
|
||||
"bind_user_name": None,
|
||||
"device_type": row.get("board"),
|
||||
"board": row.get("board"),
|
||||
"id": row.get("id"),
|
||||
"mac_address": row.get("mac_address"),
|
||||
"alias": row.get("alias"),
|
||||
"ota_upgrade": None,
|
||||
"recent_chat_time": None,
|
||||
"last_connected_at_timestamp": DeviceService._timestamp(row.get("last_connected_at")),
|
||||
"create_date_timestamp": DeviceService._timestamp(row.get("create_date")),
|
||||
# UserShowDeviceListVO pins only this field to UTC. The companion
|
||||
# epoch value still uses the configured Asia/Shanghai instant.
|
||||
"create_date": DeviceService._utc_datetime(row.get("create_date")),
|
||||
}
|
||||
|
||||
@staticmethod
|
||||
def _utc_datetime(value: Any) -> Any:
|
||||
if not isinstance(value, datetime):
|
||||
return value
|
||||
localized = value.replace(tzinfo=ZoneInfo(get_settings().timezone)) if value.tzinfo is None else value
|
||||
return localized.astimezone(timezone.utc).replace(tzinfo=None)
|
||||
|
||||
@staticmethod
|
||||
def _ota_response_record(row: Mapping[str, Any]) -> dict[str, Any]:
|
||||
result = dict(row)
|
||||
if result.get("size") is not None:
|
||||
result["size"] = str(result["size"])
|
||||
return result
|
||||
|
||||
@staticmethod
|
||||
def _timestamp(value: Any) -> int | None:
|
||||
if not isinstance(value, datetime):
|
||||
return None
|
||||
localized = value.replace(tzinfo=ZoneInfo(get_settings().timezone)) if value.tzinfo is None else value
|
||||
return int(localized.timestamp() * 1000)
|
||||
@@ -0,0 +1,21 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from functools import lru_cache
|
||||
from pathlib import Path
|
||||
|
||||
from app.core.config import get_settings
|
||||
from app.core.i18n import LANGUAGE_FILES, _load_properties, resolve_language
|
||||
|
||||
|
||||
@lru_cache(maxsize=16)
|
||||
def _validation_messages(language: str, directory: str) -> dict[str, str]:
|
||||
root = Path(directory)
|
||||
values = _load_properties(root / "validation.properties")
|
||||
localized = LANGUAGE_FILES[language].replace("messages_", "validation_")
|
||||
values.update(_load_properties(root / localized))
|
||||
return values
|
||||
|
||||
|
||||
def validation_message(key: str, accept_language: str | None) -> str:
|
||||
language = resolve_language(accept_language)
|
||||
return _validation_messages(language, str(get_settings().i18n_dir)).get(key, key)
|
||||
@@ -0,0 +1,723 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
from collections import defaultdict
|
||||
from datetime import datetime
|
||||
from typing import Any
|
||||
from zoneinfo import ZoneInfo
|
||||
|
||||
from fastapi import UploadFile
|
||||
|
||||
from app.core.config import get_settings
|
||||
from app.core.errors import AppError
|
||||
from app.core.i18n import message_for
|
||||
from app.core.redis import get_redis
|
||||
from app.core.security import AuthUser, shanghai_now_naive
|
||||
from app.core.serialization import preserve_java_map_keys
|
||||
from app.integrations.ragflow import RAGFlowClient
|
||||
from app.repositories.knowledge import KnowledgeRepository
|
||||
from app.schemas.knowledge import KnowledgeBaseBody, RetrievalBody
|
||||
|
||||
|
||||
def dataset_dto(row: dict[str, Any]) -> dict[str, Any]:
|
||||
return {
|
||||
"id": row.get("id"),
|
||||
"datasetId": row.get("dataset_id"),
|
||||
"ragModelId": row.get("rag_model_id"),
|
||||
"name": row.get("name"),
|
||||
"avatar": row.get("avatar"),
|
||||
"description": row.get("description"),
|
||||
"embeddingModel": row.get("embedding_model"),
|
||||
"permission": row.get("permission"),
|
||||
"chunkMethod": row.get("chunk_method"),
|
||||
"parserConfig": row.get("parser_config"),
|
||||
"chunkCount": None if row.get("chunk_count") is None else str(row["chunk_count"]),
|
||||
"tokenNum": None if row.get("token_num") is None else str(row["token_num"]),
|
||||
"status": row.get("status"),
|
||||
"creator": row.get("creator"),
|
||||
"createdAt": row.get("created_at"),
|
||||
"updater": row.get("updater"),
|
||||
"updatedAt": row.get("updated_at"),
|
||||
# KnowledgeBaseEntity.documentCount is Long while KnowledgeBaseDTO uses
|
||||
# Integer. Spring BeanUtils does not coerce that property, so local DTO
|
||||
# conversion leaves it null; list enrichment fills it from RAGFlow.
|
||||
"documentCount": None,
|
||||
"errorMessage": row.get("error_message"),
|
||||
}
|
||||
|
||||
|
||||
def document_dto(row: dict[str, Any]) -> dict[str, Any]:
|
||||
return {
|
||||
"id": row.get("document_id"),
|
||||
"documentId": row.get("document_id"),
|
||||
"datasetId": row.get("dataset_id"),
|
||||
"name": row.get("name"),
|
||||
# RAGFlowAdapter.mapToKnowledgeFilesDTO does not populate these two
|
||||
# fields for the immediate upload response.
|
||||
"fileType": None,
|
||||
"fileSize": row.get("size"),
|
||||
"filePath": None,
|
||||
"progress": row.get("progress"),
|
||||
"thumbnail": row.get("thumbnail"),
|
||||
"processDuration": row.get("process_duration"),
|
||||
"sourceType": row.get("source_type"),
|
||||
"metaFields": _json_object(row.get("meta_fields")),
|
||||
"chunkMethod": row.get("chunk_method"),
|
||||
"parserConfig": _json_object(row.get("parser_config")),
|
||||
"status": row.get("status"),
|
||||
"run": row.get("run"),
|
||||
"creator": row.get("creator"),
|
||||
"createdAt": row.get("created_at"),
|
||||
"updater": None,
|
||||
"updatedAt": row.get("updated_at"),
|
||||
"chunkCount": row.get("chunk_count"),
|
||||
"tokenCount": row.get("token_count"),
|
||||
"error": row.get("error"),
|
||||
"parseStatusCode": _parse_status(row.get("run")),
|
||||
}
|
||||
|
||||
|
||||
def remote_document_dto(row: dict[str, Any], dataset_id: str) -> dict[str, Any]:
|
||||
run = row.get("run")
|
||||
return {
|
||||
"id": row.get("id"),
|
||||
"documentId": row.get("id"),
|
||||
"datasetId": row.get("dataset_id") or dataset_id,
|
||||
"name": row.get("name"),
|
||||
"fileType": row.get("type"),
|
||||
"fileSize": row.get("size"),
|
||||
"filePath": None,
|
||||
"progress": row.get("progress"),
|
||||
"thumbnail": row.get("thumbnail"),
|
||||
"processDuration": row.get("process_duration"),
|
||||
"sourceType": row.get("source_type"),
|
||||
"metaFields": row.get("meta_fields"),
|
||||
"chunkMethod": row.get("chunk_method"),
|
||||
"parserConfig": row.get("parser_config"),
|
||||
"status": _remote_status(row.get("status")),
|
||||
"run": run,
|
||||
"creator": None,
|
||||
"createdAt": _millis(row.get("create_time")),
|
||||
"updater": None,
|
||||
"updatedAt": _millis(row.get("update_time")),
|
||||
"chunkCount": row.get("chunk_count") or 0,
|
||||
"tokenCount": row.get("token_count"),
|
||||
"error": row.get("progress_msg"),
|
||||
"parseStatusCode": _parse_status(run),
|
||||
}
|
||||
|
||||
|
||||
def _parse_status(run: Any) -> int:
|
||||
return {"RUNNING": 1, "CANCEL": 2, "DONE": 3, "FAIL": 4}.get(str(run or "").upper(), 0)
|
||||
|
||||
|
||||
def _json_object(value: Any) -> dict[str, Any] | None:
|
||||
if value is None:
|
||||
return None
|
||||
if isinstance(value, dict):
|
||||
return dict(value)
|
||||
try:
|
||||
parsed = json.loads(value.decode() if isinstance(value, bytes) else str(value))
|
||||
return dict(parsed) if isinstance(parsed, dict) else None
|
||||
except (ValueError, TypeError):
|
||||
return None
|
||||
|
||||
|
||||
def _millis(value: Any) -> Any:
|
||||
try:
|
||||
if value is None:
|
||||
return None
|
||||
timezone = ZoneInfo(get_settings().timezone)
|
||||
return datetime.fromtimestamp(float(value) / 1000, timezone).replace(tzinfo=None)
|
||||
except (TypeError, ValueError, OSError):
|
||||
return None
|
||||
|
||||
|
||||
def _is_blank(value: str | None) -> bool:
|
||||
return value is None or not value.strip()
|
||||
|
||||
|
||||
def _remote_status(value: Any) -> str:
|
||||
if value is None or (isinstance(value, str) and not value.strip()):
|
||||
return "1"
|
||||
return str(value)
|
||||
|
||||
|
||||
class KnowledgeBaseService:
|
||||
def __init__(self, repository: KnowledgeRepository):
|
||||
self.repository = repository
|
||||
|
||||
async def get_owned(self, identifier: str, user: AuthUser) -> dict[str, Any]:
|
||||
if not identifier.strip():
|
||||
raise AppError(10003)
|
||||
row = await self.repository.get_dataset(identifier)
|
||||
if row is None:
|
||||
raise AppError(10163)
|
||||
if row.get("creator") is None or int(row["creator"]) != user.id:
|
||||
raise AppError(10169)
|
||||
return row
|
||||
|
||||
async def page(
|
||||
self,
|
||||
user: AuthUser,
|
||||
name: str | None,
|
||||
page: int,
|
||||
page_size: int,
|
||||
language: str | None = None,
|
||||
) -> dict[str, Any]:
|
||||
rows, total = await self.repository.dataset_page(
|
||||
user.id, name, (max(page, 1) - 1) * page_size, page_size
|
||||
)
|
||||
results: list[dict[str, Any]] = []
|
||||
changed = False
|
||||
for row in rows:
|
||||
dto = dataset_dto(row)
|
||||
if row.get("dataset_id") and row.get("rag_model_id"):
|
||||
try:
|
||||
client = await self._client(str(row["rag_model_id"]))
|
||||
remote = await client.dataset_info(str(row["dataset_id"]))
|
||||
if remote is None:
|
||||
await self.repository.execute(
|
||||
"DELETE FROM ai_rag_knowledge_document WHERE dataset_id=:dataset_id",
|
||||
{"dataset_id": row["dataset_id"]},
|
||||
)
|
||||
await self.repository.delete_dataset_local(row)
|
||||
await _delete_cache_ignoring_errors(f"knowledge:base:{row['id']}")
|
||||
changed = True
|
||||
continue
|
||||
remote_name = remote.get("name")
|
||||
local_name = (
|
||||
str(remote_name).split("_", 1)[1]
|
||||
if remote_name and "_" in str(remote_name)
|
||||
else remote_name
|
||||
)
|
||||
updates: dict[str, Any] = {}
|
||||
if local_name and local_name != row.get("name"):
|
||||
updates["name"] = local_name
|
||||
dto["name"] = local_name
|
||||
if remote.get("description") != row.get("description"):
|
||||
updates["description"] = remote.get("description")
|
||||
dto["description"] = remote.get("description")
|
||||
if updates:
|
||||
await self.repository.execute(
|
||||
"UPDATE ai_rag_dataset SET name=COALESCE(:name,name),description=:description WHERE id=:id",
|
||||
{
|
||||
"name": updates.get("name"),
|
||||
"description": updates.get("description", row.get("description")),
|
||||
"id": row["id"],
|
||||
},
|
||||
)
|
||||
changed = True
|
||||
if remote.get("document_count") is not None:
|
||||
dto["documentCount"] = int(remote["document_count"])
|
||||
except Exception as exc:
|
||||
dto["documentCount"] = 0
|
||||
dto["errorMessage"] = (
|
||||
message_for(exc.code, language, *exc.params)
|
||||
if isinstance(exc, AppError)
|
||||
else str(exc)
|
||||
)
|
||||
results.append(dto)
|
||||
if changed:
|
||||
await self.repository.session.commit()
|
||||
return {"total": total, "list": results}
|
||||
|
||||
async def create(self, body: KnowledgeBaseBody, user: AuthUser) -> dict[str, Any]:
|
||||
if not _is_blank(body.name) and await self.repository.duplicate_dataset_name(user.id, str(body.name)):
|
||||
raise AppError(10170)
|
||||
rag_model_id = body.rag_model_id
|
||||
if _is_blank(rag_model_id):
|
||||
models = await self.repository.rag_models()
|
||||
if not models:
|
||||
raise AppError(10164, params=("未指定且无可用默认 RAG 模型",))
|
||||
rag_model_id = str(models[0]["id"])
|
||||
client = await self._client(str(rag_model_id))
|
||||
create_body = {
|
||||
"name": f"{user.username}_{'null' if body.name is None else body.name}",
|
||||
"avatar": body.avatar,
|
||||
"description": body.description,
|
||||
"embedding_model": body.embedding_model,
|
||||
"permission": body.permission,
|
||||
"chunk_method": body.chunk_method,
|
||||
# KnowledgeBaseDTO.parserConfig is a String, while CreateReq uses
|
||||
# ParserConfig. BeanUtils skips the incompatible property.
|
||||
"parser_config": None,
|
||||
}
|
||||
remote = await client.create_dataset(create_body)
|
||||
dataset_id = str(remote["id"])
|
||||
now = shanghai_now_naive()
|
||||
created_at = body.created_at or now
|
||||
updated_at = body.updated_at or now
|
||||
values = {
|
||||
"id": dataset_id,
|
||||
"dataset_id": dataset_id,
|
||||
"rag_model_id": rag_model_id,
|
||||
"tenant_id": remote.get("tenant_id"),
|
||||
"name": body.name,
|
||||
"avatar": remote.get("avatar") if _is_blank(body.avatar) else body.avatar,
|
||||
"description": body.description,
|
||||
"embedding_model": remote.get("embedding_model"),
|
||||
"permission": remote.get("permission"),
|
||||
"chunk_method": remote.get("chunk_method"),
|
||||
"parser_config": json.dumps(
|
||||
remote.get("parser_config"), ensure_ascii=False, separators=(",", ":")
|
||||
)
|
||||
if remote.get("parser_config") is not None
|
||||
else None,
|
||||
"chunk_count": remote.get("chunk_count") or 0,
|
||||
"document_count": remote.get("document_count") or 0,
|
||||
"token_num": remote.get("token_num") or 0,
|
||||
"status": 1,
|
||||
"creator": user.id,
|
||||
"updater": user.id,
|
||||
"created_at": created_at,
|
||||
"updated_at": updated_at,
|
||||
}
|
||||
try:
|
||||
await self.repository.insert_dataset(values)
|
||||
await self.repository.session.commit()
|
||||
except Exception as exc:
|
||||
await self.repository.session.rollback()
|
||||
try:
|
||||
await client.delete_datasets([dataset_id])
|
||||
except AppError:
|
||||
pass
|
||||
if isinstance(exc, AppError):
|
||||
raise
|
||||
raise AppError(10167, params=(f"创建知识库失败: {exc}",)) from exc
|
||||
return dataset_dto(values)
|
||||
|
||||
async def update(
|
||||
self, identifier: str, body: KnowledgeBaseBody, user: AuthUser
|
||||
) -> dict[str, Any]:
|
||||
existing = await self.get_owned(identifier, user)
|
||||
if not _is_blank(body.name) and await self.repository.duplicate_dataset_name(
|
||||
user.id, str(body.name), str(existing["id"])
|
||||
):
|
||||
raise AppError(10170)
|
||||
if not _is_blank(identifier) and await self.repository.dataset_id_conflict(
|
||||
identifier, str(existing["id"])
|
||||
):
|
||||
raise AppError(10002)
|
||||
rag_model_id = body.rag_model_id
|
||||
effective_permission = body.permission
|
||||
effective_chunk_method = body.chunk_method
|
||||
if existing.get("dataset_id") and not _is_blank(rag_model_id):
|
||||
if _is_blank(effective_permission):
|
||||
effective_permission = existing.get("permission")
|
||||
if _is_blank(effective_chunk_method):
|
||||
effective_chunk_method = existing.get("chunk_method")
|
||||
client = await self._client(str(rag_model_id))
|
||||
remote_body = {
|
||||
"name": f"{user.username}_{body.name}" if not _is_blank(body.name) else None,
|
||||
"avatar": body.avatar,
|
||||
"description": body.description,
|
||||
"embedding_model": body.embedding_model,
|
||||
"permission": effective_permission,
|
||||
"chunk_method": effective_chunk_method,
|
||||
"parser_config": _json_object(body.parser_config),
|
||||
}
|
||||
await client.update_dataset(str(existing["dataset_id"]), remote_body)
|
||||
now = shanghai_now_naive()
|
||||
updater = body.updater if body.updater is not None else user.id
|
||||
updated_at = body.updated_at or now
|
||||
values = {
|
||||
"id": existing["id"],
|
||||
# The controller injects the literal path value into datasetId,
|
||||
# even when a legacy row was found through its local primary key.
|
||||
"dataset_id": identifier,
|
||||
"rag_model_id": rag_model_id,
|
||||
"name": body.name,
|
||||
"avatar": body.avatar,
|
||||
"description": body.description,
|
||||
"embedding_model": body.embedding_model,
|
||||
"permission": effective_permission,
|
||||
"chunk_method": effective_chunk_method,
|
||||
"parser_config": body.parser_config,
|
||||
"chunk_count": body.chunk_count,
|
||||
"token_num": body.token_num,
|
||||
"status": body.status,
|
||||
"creator": body.creator,
|
||||
"created_at": body.created_at,
|
||||
"updater": updater,
|
||||
"updated_at": updated_at,
|
||||
}
|
||||
try:
|
||||
await self.repository.update_dataset(values)
|
||||
# Java performs cache eviction inside the database transaction;
|
||||
# an eviction failure therefore rolls this update back.
|
||||
await get_redis().delete(f"knowledge:base:{existing['id']}")
|
||||
await self.repository.session.commit()
|
||||
except Exception:
|
||||
await self.repository.session.rollback()
|
||||
raise
|
||||
# BeanUtils copies request nulls onto the in-memory entity before
|
||||
# MyBatis' NOT_NULL update strategy preserves the stored columns. The
|
||||
# Java response is built from that in-memory entity, so its null fields
|
||||
# intentionally differ from a subsequent GET of the row.
|
||||
return dataset_dto(values)
|
||||
|
||||
async def delete(self, identifier: str, user: AuthUser, language: str | None = None) -> None:
|
||||
row = await self.get_owned(identifier, user)
|
||||
documents = await self.repository.all_documents(str(row["dataset_id"]))
|
||||
if documents:
|
||||
# Java's document orchestration necessarily resolves the adapter
|
||||
# when child records exist.
|
||||
client = await self._client(str(row.get("rag_model_id") or ""))
|
||||
ids = [str(item["document_id"]) for item in documents]
|
||||
if any(item.get("run") == "RUNNING" for item in documents):
|
||||
raise AppError(10199)
|
||||
try:
|
||||
await client.delete_documents(str(row["dataset_id"]), ids)
|
||||
except Exception as exc:
|
||||
raise _document_delete_error(exc, language) from exc
|
||||
await self.repository.delete_document_shadows(str(row["dataset_id"]), ids)
|
||||
await self.repository.update_stats(
|
||||
str(row["dataset_id"]),
|
||||
-len(ids),
|
||||
-sum(int(item.get("chunk_count") or 0) for item in documents),
|
||||
-sum(int(item.get("token_count") or 0) for item in documents),
|
||||
)
|
||||
# deleteDocuments is NOT_SUPPORTED in Java and its shadow cleanup
|
||||
# commits before the outer dataset transaction continues.
|
||||
await self.repository.session.commit()
|
||||
await _delete_cache_ignoring_errors(f"knowledge:base:{row['dataset_id']}")
|
||||
if not _is_blank(row.get("rag_model_id")) and not _is_blank(row.get("dataset_id")):
|
||||
client = await self._client(str(row["rag_model_id"]))
|
||||
await client.delete_datasets([str(row["dataset_id"])])
|
||||
await self.repository.delete_dataset_local(row)
|
||||
try:
|
||||
await get_redis().delete(f"knowledge:base:{row['id']}")
|
||||
await self.repository.session.commit()
|
||||
except Exception:
|
||||
await self.repository.session.rollback()
|
||||
raise
|
||||
|
||||
async def batch_delete(
|
||||
self, identifiers: list[str], user: AuthUser, language: str | None = None
|
||||
) -> None:
|
||||
rows = await self.repository.datasets_by_ids(identifiers)
|
||||
for row in rows:
|
||||
if row.get("creator") is None or int(row["creator"]) != user.id:
|
||||
raise AppError(10169)
|
||||
# Preserve Java's sequential external calls and stop-on-first-error semantics.
|
||||
for row in rows:
|
||||
await self.delete(str(row["dataset_id"]), user, language)
|
||||
|
||||
async def rag_models(self) -> list[dict[str, Any]]:
|
||||
rows = await self.repository.rag_models()
|
||||
result: list[dict[str, Any]] = []
|
||||
for row in rows:
|
||||
result.append(
|
||||
{
|
||||
"id": row.get("id"),
|
||||
"modelType": None,
|
||||
"modelCode": None,
|
||||
"modelName": row.get("model_name"),
|
||||
"isDefault": None,
|
||||
"isEnabled": None,
|
||||
# ModelConfigEntity.configJson is a JSONObject. Jackson
|
||||
# preserves its dynamic snake_case keys instead of applying
|
||||
# the DTO property naming strategy recursively.
|
||||
"configJson": preserve_java_map_keys(_json_object(row.get("config_json"))),
|
||||
"docLink": None,
|
||||
"remark": None,
|
||||
"sort": None,
|
||||
"updater": None,
|
||||
"updateDate": None,
|
||||
"creator": None,
|
||||
"createDate": None,
|
||||
}
|
||||
)
|
||||
return result
|
||||
|
||||
async def _client(self, model_id: str) -> RAGFlowClient:
|
||||
config = await self.repository.rag_config(model_id)
|
||||
adapter_type = config.get("type")
|
||||
if adapter_type != "ragflow":
|
||||
raise AppError(10184, params=(f"适配器类型未注册: {adapter_type}",))
|
||||
try:
|
||||
return RAGFlowClient(config)
|
||||
except AppError as exc:
|
||||
# KnowledgeBaseAdapterFactory wraps adapter initialization and
|
||||
# validateConfig failures as RAG_ADAPTER_CREATION_FAILED.
|
||||
if exc.code in {10171, 10172, 10173, 10174}:
|
||||
raise AppError(10186) from exc
|
||||
raise
|
||||
|
||||
|
||||
class KnowledgeDocumentService:
|
||||
def __init__(self, repository: KnowledgeRepository):
|
||||
self.repository = repository
|
||||
self.datasets = KnowledgeBaseService(repository)
|
||||
|
||||
async def page(
|
||||
self,
|
||||
dataset_id: str,
|
||||
user: AuthUser,
|
||||
*,
|
||||
name: str | None,
|
||||
status: str | None,
|
||||
page: int,
|
||||
page_size: int,
|
||||
) -> dict[str, Any]:
|
||||
await self.datasets.get_owned(dataset_id, user)
|
||||
try:
|
||||
await self.reconcile(dataset_id, creator=user.id)
|
||||
except Exception:
|
||||
await self.repository.session.rollback()
|
||||
rows, total = await self.repository.documents_page(
|
||||
dataset_id,
|
||||
name=name,
|
||||
status=status,
|
||||
offset=(max(page, 1) - 1) * page_size,
|
||||
limit=page_size,
|
||||
)
|
||||
return {"total": total, "list": [document_dto(row) for row in rows]}
|
||||
|
||||
async def upload(
|
||||
self,
|
||||
dataset_id: str,
|
||||
user: AuthUser,
|
||||
file: UploadFile,
|
||||
*,
|
||||
name: str | None,
|
||||
meta_fields: dict[str, Any] | None,
|
||||
chunk_method: str | None,
|
||||
parser_config: dict[str, Any] | None,
|
||||
) -> dict[str, Any]:
|
||||
await self.datasets.get_owned(dataset_id, user)
|
||||
content = await file.read()
|
||||
if not dataset_id.strip() or not content:
|
||||
raise AppError(10003)
|
||||
file_name = file.filename if _is_blank(name) else name
|
||||
if _is_blank(file_name):
|
||||
raise AppError(10179)
|
||||
assert file_name is not None
|
||||
client = await self._client_for_dataset(dataset_id)
|
||||
remote = await client.upload_document(
|
||||
dataset_id,
|
||||
file,
|
||||
content,
|
||||
name=file_name,
|
||||
meta_fields=meta_fields,
|
||||
chunk_method=chunk_method,
|
||||
parser_config=parser_config,
|
||||
)
|
||||
if not remote.get("id"):
|
||||
raise AppError(10167, params=("远程上传成功但未返回有效 DocumentID",))
|
||||
remote.setdefault("dataset_id", dataset_id)
|
||||
shadow = dict(remote)
|
||||
if _is_blank(str(shadow.get("name")) if shadow.get("name") is not None else None):
|
||||
shadow["name"] = file_name
|
||||
# Java stores the original controller values in the shadow row, even
|
||||
# when invalid chunk methods were omitted from the RAGFlow request.
|
||||
shadow["chunk_method"] = chunk_method
|
||||
shadow["parser_config"] = parser_config
|
||||
inserted = await self.repository.upsert_document(dataset_id, shadow, creator=user.id)
|
||||
if inserted:
|
||||
await self.repository.update_stats(dataset_id, 1, 0, 0)
|
||||
await self.repository.session.commit()
|
||||
return remote_document_dto(remote, dataset_id)
|
||||
|
||||
async def delete(
|
||||
self,
|
||||
dataset_id: str,
|
||||
ids: list[str] | None,
|
||||
user: AuthUser,
|
||||
language: str | None = None,
|
||||
) -> None:
|
||||
await self.datasets.get_owned(dataset_id, user)
|
||||
if not ids:
|
||||
raise AppError(10178)
|
||||
rows = await self.repository.documents_by_remote_ids(dataset_id, ids)
|
||||
if len(rows) != len(ids):
|
||||
raise AppError(10169)
|
||||
if any(row.get("run") == "RUNNING" for row in rows):
|
||||
raise AppError(10199)
|
||||
chunks = sum(int(row.get("chunk_count") or 0) for row in rows)
|
||||
tokens = sum(int(row.get("token_count") or 0) for row in rows)
|
||||
client = await self._client_for_dataset(dataset_id)
|
||||
try:
|
||||
await client.delete_documents(dataset_id, ids)
|
||||
except Exception as exc:
|
||||
raise _document_delete_error(exc, language) from exc
|
||||
deleted = await self.repository.delete_document_shadows(dataset_id, ids)
|
||||
if deleted:
|
||||
await self.repository.update_stats(dataset_id, -len(ids), -chunks, -tokens)
|
||||
await self.repository.session.commit()
|
||||
await _delete_cache_ignoring_errors(f"knowledge:base:{dataset_id}")
|
||||
|
||||
async def parse(self, dataset_id: str, ids: list[str], user: AuthUser) -> bool:
|
||||
await self.datasets.get_owned(dataset_id, user)
|
||||
if not ids:
|
||||
raise AppError(10178)
|
||||
client = await self._client_for_dataset(dataset_id)
|
||||
await client.parse_documents(dataset_id, ids)
|
||||
await self.repository.mark_documents_running(dataset_id, ids, shanghai_now_naive())
|
||||
await self.repository.session.commit()
|
||||
return True
|
||||
|
||||
async def chunks(
|
||||
self,
|
||||
dataset_id: str,
|
||||
document_id: str,
|
||||
user: AuthUser,
|
||||
*,
|
||||
page: int,
|
||||
page_size: int,
|
||||
keywords: str | None,
|
||||
chunk_id: str | None,
|
||||
) -> dict[str, Any]:
|
||||
await self.datasets.get_owned(dataset_id, user)
|
||||
client = await self._client_for_dataset(dataset_id)
|
||||
return await client.chunks(
|
||||
dataset_id,
|
||||
document_id,
|
||||
{"page": page, "page_size": page_size, "keywords": keywords, "id": chunk_id},
|
||||
)
|
||||
|
||||
async def retrieval(self, dataset_id: str, body: RetrievalBody, user: AuthUser) -> dict[str, Any]:
|
||||
await self.datasets.get_owned(dataset_id, user)
|
||||
dataset_ids = body.dataset_ids or [dataset_id]
|
||||
if not dataset_ids:
|
||||
raise AppError(500, "未指定召回测试的知识库")
|
||||
page = body.page if body.page is not None and body.page >= 1 else 1
|
||||
page_size = body.page_size if body.page_size is not None and body.page_size >= 1 else 100
|
||||
top_k = body.top_k if body.top_k is None or body.top_k >= 1 else 1024
|
||||
threshold = body.similarity_threshold
|
||||
if threshold is not None:
|
||||
threshold = 0.2 if threshold < 0 else min(threshold, 1.0)
|
||||
payload: dict[str, Any] = {
|
||||
"dataset_ids": dataset_ids,
|
||||
"document_ids": body.document_ids,
|
||||
"question": body.question,
|
||||
"page": page,
|
||||
"page_size": page_size,
|
||||
"similarity_threshold": threshold,
|
||||
"vector_similarity_weight": body.vector_similarity_weight,
|
||||
"top_k": top_k,
|
||||
"rerank_id": body.rerank_id,
|
||||
"highlight": body.highlight,
|
||||
"keyword": body.keyword,
|
||||
"cross_languages": body.cross_languages,
|
||||
"metadata_condition": body.metadata_condition,
|
||||
}
|
||||
payload = {key: value for key, value in payload.items() if value is not None}
|
||||
client = await self._client_for_dataset(dataset_ids[0])
|
||||
return await client.retrieval(payload)
|
||||
|
||||
async def reconcile(self, dataset_id: str, *, creator: int | None = None) -> int:
|
||||
client = await self._client_for_dataset(dataset_id)
|
||||
remote: list[dict[str, Any]] = []
|
||||
page, total = 1, 2**63 - 1
|
||||
while (page - 1) * 100 < total:
|
||||
rows, total = await client.documents(dataset_id, page=page, page_size=100)
|
||||
if not rows:
|
||||
break
|
||||
remote.extend(rows)
|
||||
page += 1
|
||||
local = await self.repository.all_documents(dataset_id)
|
||||
remote_map = {str(item.get("id")): item for item in remote if item.get("id")}
|
||||
local_map = {str(item["document_id"]): item for item in local}
|
||||
new_count = 0
|
||||
for document_id, item in remote_map.items():
|
||||
prior = local_map.get(document_id)
|
||||
inserted = await self.repository.upsert_document(dataset_id, item, creator=creator)
|
||||
if inserted:
|
||||
new_count += 1
|
||||
await self.repository.update_stats(
|
||||
dataset_id, 1, int(item.get("chunk_count") or 0), int(item.get("token_count") or 0)
|
||||
)
|
||||
elif prior:
|
||||
await self.repository.update_stats(
|
||||
dataset_id,
|
||||
0,
|
||||
int(item.get("chunk_count") or 0) - int(prior.get("chunk_count") or 0),
|
||||
int(item.get("token_count") or 0) - int(prior.get("token_count") or 0),
|
||||
)
|
||||
deleted_ids = [identifier for identifier in local_map if identifier not in remote_map]
|
||||
if deleted_ids:
|
||||
deleted_rows = [local_map[identifier] for identifier in deleted_ids]
|
||||
await self.repository.delete_document_shadows(dataset_id, deleted_ids)
|
||||
await self.repository.update_stats(
|
||||
dataset_id,
|
||||
-len(deleted_ids),
|
||||
-sum(int(row.get("chunk_count") or 0) for row in deleted_rows),
|
||||
-sum(int(row.get("token_count") or 0) for row in deleted_rows),
|
||||
)
|
||||
await self.repository.session.commit()
|
||||
return new_count
|
||||
|
||||
async def sync_running(self) -> int:
|
||||
rows = await self.repository.running_documents()
|
||||
grouped: defaultdict[str, list[dict[str, Any]]] = defaultdict(list)
|
||||
for row in rows:
|
||||
grouped[str(row["dataset_id"])].append(row)
|
||||
updates = 0
|
||||
for dataset_id, documents in grouped.items():
|
||||
try:
|
||||
client = await self._client_for_dataset(dataset_id)
|
||||
except Exception:
|
||||
await self.repository.session.rollback()
|
||||
continue
|
||||
for local in documents:
|
||||
try:
|
||||
remote, _ = await client.documents(
|
||||
dataset_id, page=1, page_size=1, document_id=str(local["document_id"])
|
||||
)
|
||||
if not remote:
|
||||
await self.repository.mark_document_remote_deleted(
|
||||
str(local["document_id"]), shanghai_now_naive()
|
||||
)
|
||||
await self.repository.session.commit()
|
||||
updates += 1
|
||||
continue
|
||||
remote_status = remote[0].get("status")
|
||||
remote_run = remote[0].get("run")
|
||||
status_changed = remote_status is not None and str(remote_status) != str(local.get("status"))
|
||||
run_changed = remote_run is not None and str(remote_run) != str(local.get("run"))
|
||||
is_processing = remote_run in {"RUNNING", "UNSTART"}
|
||||
if not (status_changed or run_changed or is_processing):
|
||||
await self.repository.session.commit()
|
||||
continue
|
||||
before_tokens = int(local.get("token_count") or 0)
|
||||
await self.repository.sync_running_document(
|
||||
dataset_id,
|
||||
str(local["document_id"]),
|
||||
remote[0],
|
||||
shanghai_now_naive(),
|
||||
)
|
||||
delta = int(remote[0].get("token_count") or 0) - before_tokens
|
||||
if delta:
|
||||
await self.repository.update_stats(dataset_id, 0, 0, delta)
|
||||
await self.repository.session.commit()
|
||||
updates += 1
|
||||
except Exception:
|
||||
await self.repository.session.rollback()
|
||||
continue
|
||||
return updates
|
||||
|
||||
async def _client_for_dataset(self, dataset_id: str) -> RAGFlowClient:
|
||||
row = await self.repository.get_dataset(dataset_id)
|
||||
if row is None or not row.get("rag_model_id"):
|
||||
raise AppError(10164)
|
||||
return await self.datasets._client(str(row["rag_model_id"]))
|
||||
|
||||
|
||||
def _document_delete_error(exc: Exception, language: str | None) -> AppError:
|
||||
"""Match `new RenException(e.getMessage())` in the Java delete flow."""
|
||||
if isinstance(exc, AppError):
|
||||
message = exc.message or message_for(exc.code, language, *exc.params)
|
||||
else:
|
||||
message = str(exc)
|
||||
return AppError(500, message)
|
||||
|
||||
|
||||
async def _delete_cache_ignoring_errors(key: str) -> None:
|
||||
try:
|
||||
await get_redis().delete(key)
|
||||
except Exception:
|
||||
# The Java document cleanup and remote-missing cleanup explicitly log
|
||||
# and continue when Redis is unavailable.
|
||||
return
|
||||
@@ -0,0 +1,306 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import copy
|
||||
import json
|
||||
import uuid
|
||||
from typing import Any
|
||||
|
||||
from app.core.errors import AppError
|
||||
from app.core.redis import get_redis
|
||||
from app.core.security import AuthUser, shanghai_now_naive
|
||||
from app.repositories.model import ModelRepository, parse_json_object
|
||||
from app.schemas.model import ModelConfigBody, ModelProviderBody
|
||||
|
||||
SENSITIVE_FIELDS = {
|
||||
"api_key",
|
||||
"personal_access_token",
|
||||
"access_token",
|
||||
"token",
|
||||
"secret",
|
||||
"access_key_secret",
|
||||
"secret_key",
|
||||
}
|
||||
|
||||
|
||||
def _mask_middle(value: str) -> str:
|
||||
if not value.strip() or len(value) == 1:
|
||||
return value
|
||||
if len(value) <= 8:
|
||||
return value[:2] + "****" + value[-2:]
|
||||
return value[:4] + "*" * (len(value) - 8) + value[-4:]
|
||||
|
||||
|
||||
def mask_sensitive(value: Any) -> Any:
|
||||
if not isinstance(value, dict):
|
||||
return value
|
||||
result: dict[str, Any] = {}
|
||||
for key, item in value.items():
|
||||
if key.lower() in SENSITIVE_FIELDS and isinstance(item, str):
|
||||
result[key] = _mask_middle(item)
|
||||
elif isinstance(item, dict):
|
||||
result[key] = mask_sensitive(item)
|
||||
else:
|
||||
result[key] = copy.deepcopy(item)
|
||||
return result
|
||||
|
||||
|
||||
def _merge_config(original: dict[str, Any], updated: dict[str, Any]) -> dict[str, Any]:
|
||||
result = copy.deepcopy(original)
|
||||
for key, value in updated.items():
|
||||
if key.lower() in SENSITIVE_FIELDS:
|
||||
if isinstance(value, str) and "****" not in value:
|
||||
result[key] = value
|
||||
elif isinstance(value, dict):
|
||||
child = result.get(key)
|
||||
result[key] = _merge_config(child if isinstance(child, dict) else {}, value)
|
||||
else:
|
||||
result[key] = copy.deepcopy(value)
|
||||
for key in list(result):
|
||||
if key not in updated and key.lower() not in SENSITIVE_FIELDS:
|
||||
del result[key]
|
||||
return result
|
||||
|
||||
|
||||
def _model_dto(row: dict[str, Any], *, masked: bool = True) -> dict[str, Any]:
|
||||
config = parse_json_object(row.get("config_json"))
|
||||
return {
|
||||
"id": row.get("id"),
|
||||
"modelType": row.get("model_type"),
|
||||
"modelCode": row.get("model_code"),
|
||||
"modelName": row.get("model_name"),
|
||||
"isDefault": row.get("is_default"),
|
||||
"isEnabled": row.get("is_enabled"),
|
||||
"configJson": mask_sensitive(config) if masked else config,
|
||||
"docLink": row.get("doc_link"),
|
||||
"remark": row.get("remark"),
|
||||
"sort": row.get("sort"),
|
||||
}
|
||||
|
||||
|
||||
class ModelService:
|
||||
def __init__(self, repository: ModelRepository):
|
||||
self.repository = repository
|
||||
|
||||
async def names(self, model_type: str, model_name: str | None) -> list[dict[str, Any]]:
|
||||
return [
|
||||
{"id": row.get("id"), "modelName": row.get("model_name")}
|
||||
for row in await self.repository.list_model_names(model_type, model_name)
|
||||
]
|
||||
|
||||
async def llm_names(self, model_name: str | None) -> list[dict[str, Any]]:
|
||||
result: list[dict[str, Any]] = []
|
||||
for row in await self.repository.list_llm_names(model_name):
|
||||
config = parse_json_object(row.get("config_json")) or {}
|
||||
result.append(
|
||||
{"id": row.get("id"), "modelName": row.get("model_name"), "type": str(config.get("type", ""))}
|
||||
)
|
||||
return result
|
||||
|
||||
async def model_page(self, model_type: str, model_name: str | None, page: str, limit: str) -> dict[str, Any]:
|
||||
current, size = max(int(page), 1), int(limit)
|
||||
rows, total = await self.repository.list_model_configs(
|
||||
model_type=model_type,
|
||||
model_name=model_name,
|
||||
offset=(current - 1) * size,
|
||||
limit=size,
|
||||
)
|
||||
return {"total": total, "list": [_model_dto(row) for row in rows]}
|
||||
|
||||
async def get_model(self, model_id: str) -> dict[str, Any] | None:
|
||||
row = await self.repository.get_model(model_id)
|
||||
return _model_dto(row) if row else None
|
||||
|
||||
async def add(self, model_type: str, provider_code: str, body: ModelConfigBody) -> dict[str, Any]:
|
||||
if not model_type.strip() or not provider_code.strip():
|
||||
raise AppError(10131)
|
||||
model_id = body.id or uuid.uuid4().hex
|
||||
values = {
|
||||
"id": model_id,
|
||||
"model_type": model_type,
|
||||
"model_code": body.model_code,
|
||||
"model_name": body.model_name,
|
||||
"is_default": 0,
|
||||
"is_enabled": body.is_enabled,
|
||||
"config_json": json.dumps(body.config_json, ensure_ascii=False) if body.config_json is not None else None,
|
||||
"doc_link": body.doc_link,
|
||||
"remark": body.remark,
|
||||
"sort": body.sort,
|
||||
}
|
||||
async with self.repository.session.begin():
|
||||
# Keep the read and write in one transaction. A query before
|
||||
# ``begin()`` triggers SQLAlchemy autobegin and makes the explicit
|
||||
# transaction fail with InvalidRequestError.
|
||||
if await self.repository.get_provider(model_type, provider_code) is None:
|
||||
raise AppError(10162)
|
||||
await self.repository.insert_model(values)
|
||||
return _model_dto(values)
|
||||
|
||||
async def edit(
|
||||
self, model_type: str, provider_code: str, model_id: str, body: ModelConfigBody
|
||||
) -> dict[str, Any]:
|
||||
if not model_type.strip() or not provider_code.strip():
|
||||
raise AppError(10131)
|
||||
async with self.repository.session.begin():
|
||||
if await self.repository.get_provider(model_type, provider_code) is None:
|
||||
raise AppError(10162)
|
||||
original = await self.repository.get_model(model_id, for_update=True)
|
||||
if original is None:
|
||||
raise AppError(10051)
|
||||
updated_config = body.config_json
|
||||
if updated_config is not None and "llm" in updated_config:
|
||||
llm = await self.repository.get_model(str(updated_config["llm"]))
|
||||
llm_config = parse_json_object(llm.get("config_json")) if llm else None
|
||||
if llm is None or str(llm.get("model_type") or "").upper() != "LLM":
|
||||
raise AppError(10092)
|
||||
if llm_config and "type" in llm_config and llm_config["type"] not in {"openai", "ollama"}:
|
||||
raise AppError(10049)
|
||||
original_config = parse_json_object(original.get("config_json"))
|
||||
merged = (
|
||||
_merge_config(original_config, updated_config)
|
||||
if original_config is not None and updated_config is not None
|
||||
else original_config
|
||||
)
|
||||
values = {
|
||||
"id": model_id,
|
||||
"model_type": model_type,
|
||||
"model_code": original.get("model_code"),
|
||||
"model_name": body.model_name,
|
||||
"is_default": original.get("is_default"),
|
||||
"is_enabled": body.is_enabled,
|
||||
"config_json": json.dumps(merged, ensure_ascii=False) if merged is not None else None,
|
||||
"doc_link": original.get("doc_link"),
|
||||
"remark": body.remark,
|
||||
"sort": body.sort,
|
||||
}
|
||||
await self.repository.update_model(values)
|
||||
await self._clear_cache(model_id)
|
||||
return _model_dto(values)
|
||||
|
||||
async def delete(self, model_id: str) -> None:
|
||||
if not model_id.strip():
|
||||
raise AppError(10006)
|
||||
async with self.repository.session.begin():
|
||||
model = await self.repository.get_model(model_id, for_update=True)
|
||||
if model and int(model.get("is_default") or 0) == 1:
|
||||
raise AppError(10064)
|
||||
agents = await self.repository.model_agent_references(model_id)
|
||||
if agents:
|
||||
raise AppError(10093, params=("、".join(agents),))
|
||||
if model and str(model.get("model_type") or "").upper() == "LLM":
|
||||
if await self.repository.intent_reference_count(model_id):
|
||||
raise AppError(10094)
|
||||
await self.repository.delete_model(model_id)
|
||||
await self._clear_cache(model_id)
|
||||
|
||||
async def enable(self, model_id: str, status: int) -> str | None:
|
||||
async with self.repository.session.begin():
|
||||
model = await self.repository.get_model(model_id, for_update=True)
|
||||
if model is None:
|
||||
return "模型配置不存在"
|
||||
if status == 0 and int(model.get("is_default") or 0) > 0:
|
||||
return "默认模型配置不允许关闭"
|
||||
await self.repository.set_model_enabled(model_id, status)
|
||||
await self._clear_cache(model_id)
|
||||
return None
|
||||
|
||||
async def set_default(self, model_id: str) -> str | None:
|
||||
async with self.repository.session.begin():
|
||||
model = await self.repository.get_model(model_id, for_update=True)
|
||||
if model is None:
|
||||
return "模型配置不存在"
|
||||
model_type = str(model.get("model_type") or "")
|
||||
await self.repository.set_models_default(model_type, 0)
|
||||
await self.repository.execute(
|
||||
"UPDATE ai_model_config SET is_enabled=1, is_default=1 WHERE id=:id", {"id": model_id}
|
||||
)
|
||||
await self.repository.update_default_template_models(model_type, model_id)
|
||||
await self._clear_type_cache(model_type)
|
||||
return None
|
||||
|
||||
async def _clear_cache(self, model_id: str) -> None:
|
||||
redis = get_redis()
|
||||
await redis.delete(f"model:data:{model_id}", f"model:name:{model_id}")
|
||||
|
||||
async def _clear_type_cache(self, model_type: str) -> None:
|
||||
rows = await self.repository.fetch_all(
|
||||
"SELECT id FROM ai_model_config WHERE model_type=:type", {"type": model_type}
|
||||
)
|
||||
if rows:
|
||||
redis = get_redis()
|
||||
keys = [key for row in rows for key in (f"model:data:{row['id']}", f"model:name:{row['id']}")]
|
||||
await redis.delete(*keys)
|
||||
|
||||
|
||||
class ModelProviderService:
|
||||
def __init__(self, repository: ModelRepository):
|
||||
self.repository = repository
|
||||
|
||||
async def page(
|
||||
self, model_type: str | None, name: str | None, page: str, limit: str
|
||||
) -> dict[str, Any]:
|
||||
current, size = max(int(page), 1), int(limit)
|
||||
rows, total = await self.repository.list_providers(
|
||||
model_type=model_type, name=name, offset=(current - 1) * size, limit=size
|
||||
)
|
||||
return {"total": total, "list": rows}
|
||||
|
||||
@staticmethod
|
||||
def _validate(body: ModelProviderBody, *, update: bool) -> None:
|
||||
if update and (body.id is None or not body.id.strip()):
|
||||
raise AppError(10034, "id不能为空")
|
||||
for field, message in (
|
||||
(body.provider_code, "providerCode不能为空"),
|
||||
(body.model_type, "modelType不能为空"),
|
||||
(body.name, "name不能为空"),
|
||||
(body.fields, "fields(JSON格式)不能为空"),
|
||||
):
|
||||
if field is None or not field.strip():
|
||||
raise AppError(10034, message)
|
||||
if body.sort is None:
|
||||
raise AppError(10034, "sort不能为空")
|
||||
|
||||
async def add(self, body: ModelProviderBody, user: AuthUser) -> dict[str, Any]:
|
||||
self._validate(body, update=False)
|
||||
now = shanghai_now_naive()
|
||||
values = {
|
||||
"id": body.id or uuid.uuid4().hex,
|
||||
"model_type": body.model_type,
|
||||
"provider_code": body.provider_code,
|
||||
"name": body.name,
|
||||
"fields": body.fields,
|
||||
"sort": body.sort,
|
||||
"creator": user.id,
|
||||
"updater": user.id,
|
||||
"now": now,
|
||||
}
|
||||
async with self.repository.session.begin():
|
||||
await self.repository.insert_provider(values)
|
||||
return {
|
||||
# The Java service returns the request DTO, not the entity on which
|
||||
# MyBatis-Plus generated the UUID. Therefore an omitted id remains
|
||||
# null in the response even though the stored row has an id.
|
||||
"id": body.id, "modelType": body.model_type, "providerCode": body.provider_code,
|
||||
"name": body.name, "fields": body.fields, "sort": body.sort, "creator": user.id,
|
||||
"updater": user.id, "createDate": now, "updateDate": now,
|
||||
}
|
||||
|
||||
async def edit(self, body: ModelProviderBody, user: AuthUser) -> dict[str, Any]:
|
||||
self._validate(body, update=True)
|
||||
now = shanghai_now_naive()
|
||||
values = {
|
||||
"id": body.id, "model_type": body.model_type, "provider_code": body.provider_code,
|
||||
"name": body.name, "fields": body.fields, "sort": body.sort, "updater": user.id, "now": now,
|
||||
}
|
||||
async with self.repository.session.begin():
|
||||
if await self.repository.update_provider(values) == 0:
|
||||
raise AppError(10066)
|
||||
return {
|
||||
"id": body.id, "modelType": body.model_type, "providerCode": body.provider_code,
|
||||
"name": body.name, "fields": body.fields, "sort": body.sort, "updater": user.id,
|
||||
"updateDate": now, "creator": None, "createDate": None,
|
||||
}
|
||||
|
||||
async def delete(self, ids: list[str]) -> None:
|
||||
async with self.repository.session.begin():
|
||||
if await self.repository.delete_providers(ids) == 0:
|
||||
raise AppError(10043)
|
||||
@@ -0,0 +1,486 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import base64
|
||||
import hashlib
|
||||
import hmac
|
||||
import io
|
||||
import json
|
||||
import logging
|
||||
import re
|
||||
import secrets
|
||||
import string
|
||||
import time
|
||||
import urllib.parse
|
||||
import uuid
|
||||
from datetime import datetime, timedelta
|
||||
from typing import Any, Protocol, cast
|
||||
|
||||
import httpx
|
||||
from fastapi import Request
|
||||
from PIL import Image, ImageDraw, ImageFont
|
||||
from redis.asyncio import Redis
|
||||
|
||||
from app.core.config import get_settings
|
||||
from app.core.crypto import bcrypt_hash, bcrypt_matches, generate_database_token, sm2_decrypt_c1c3c2
|
||||
from app.core.errors import AppError, ErrorCode
|
||||
from app.core.ids import snowflake
|
||||
from app.core.redis import JavaRedisCodec, get_redis
|
||||
from app.core.security import AuthUser, shanghai_now_naive
|
||||
from app.repositories.security import SecurityRepository
|
||||
from app.schemas.security import (
|
||||
LoginRequest,
|
||||
PasswordChangeRequest,
|
||||
RetrievePasswordRequest,
|
||||
SmsVerificationRequest,
|
||||
)
|
||||
from app.services.java_validation import validation_message
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
TOKEN_EXPIRE_SECONDS = 12 * 60 * 60
|
||||
CAPTCHA_TTL_SECONDS = 5 * 60
|
||||
CAPTCHA_LENGTH = 5
|
||||
PHONE_PATTERN = re.compile(r"^\+[1-9]\d{0,3}[1-9]\d{4,14}$")
|
||||
STRONG_PASSWORD = re.compile(r"^(?=.*[0-9])(?=.*[a-z])(?=.*[A-Z]).+$")
|
||||
|
||||
|
||||
class SmsSender(Protocol):
|
||||
async def send_verification_code(self, phone: str | None, code: str) -> None: ...
|
||||
|
||||
|
||||
class AliyunSmsSender:
|
||||
"""Minimal implementation of the Aliyun Dysmsapi RPC request used by the Java SDK."""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
repository: SecurityRepository,
|
||||
*,
|
||||
redis: Redis | None = None,
|
||||
client: httpx.AsyncClient | None = None,
|
||||
endpoint: str = "https://dysmsapi.aliyuncs.com/",
|
||||
):
|
||||
self.repository = repository
|
||||
self.redis = redis or get_redis()
|
||||
self.client = client
|
||||
self.endpoint = endpoint
|
||||
|
||||
async def send_verification_code(self, phone: str | None, code: str) -> None:
|
||||
access_key_id = await self._param("aliyun.sms.access_key_id") or ""
|
||||
access_key_secret = await self._param("aliyun.sms.access_key_secret") or ""
|
||||
sign_name = await self._param("aliyun.sms.sign_name") or ""
|
||||
template_code = await self._param("aliyun.sms.sms_code_template_code") or ""
|
||||
# The Tea SDK constructs its client before the refundable send block;
|
||||
# blank credentials therefore map to SMS_CONNECTION_FAILED (10056).
|
||||
if not access_key_id.strip() or not access_key_secret.strip():
|
||||
raise AppError(10056)
|
||||
timestamp = time.strftime("%Y-%m-%dT%H:%M:%SZ", time.gmtime())
|
||||
params: dict[str, str] = {
|
||||
"AccessKeyId": access_key_id,
|
||||
"Action": "SendSms",
|
||||
"Format": "JSON",
|
||||
"RegionId": "cn-hangzhou",
|
||||
"SignatureMethod": "HMAC-SHA1",
|
||||
"SignatureNonce": str(uuid.uuid4()),
|
||||
"SignatureVersion": "1.0",
|
||||
"SignName": sign_name,
|
||||
"TemplateCode": template_code,
|
||||
"TemplateParam": json.dumps({"code": code}, ensure_ascii=False, separators=(",", ":")),
|
||||
"Timestamp": timestamp,
|
||||
"Version": "2017-05-25",
|
||||
}
|
||||
if phone is not None:
|
||||
params["PhoneNumbers"] = phone
|
||||
params["Signature"] = self._signature(params, access_key_secret)
|
||||
if self.client is not None:
|
||||
response = await self.client.post(self.endpoint, data=params)
|
||||
response.raise_for_status()
|
||||
return
|
||||
timeout = get_settings().external_request_timeout_seconds
|
||||
async with httpx.AsyncClient(timeout=timeout) as client:
|
||||
response = await client.post(self.endpoint, data=params)
|
||||
response.raise_for_status()
|
||||
|
||||
async def _param(self, code: str) -> str | None:
|
||||
cached = JavaRedisCodec.decode(await cast(Any, self.redis.hget)("sys:params", code))
|
||||
if cached is not None:
|
||||
return str(cached)
|
||||
value = await self.repository.get_param_value(code)
|
||||
if value is not None:
|
||||
await cast(Any, self.redis.hset)("sys:params", code, JavaRedisCodec.encode(value))
|
||||
await cast(Any, self.redis.expire)("sys:params", 24 * 60 * 60)
|
||||
return value
|
||||
|
||||
@classmethod
|
||||
def _signature(cls, params: dict[str, str], secret: str) -> str:
|
||||
canonical = "&".join(
|
||||
f"{cls._percent_encode(key)}={cls._percent_encode(value)}" for key, value in sorted(params.items())
|
||||
)
|
||||
string_to_sign = f"POST&%2F&{cls._percent_encode(canonical)}"
|
||||
digest = hmac.new(
|
||||
f"{secret}&".encode(),
|
||||
string_to_sign.encode(),
|
||||
digestmod=hashlib.sha1, # noqa: S324 - mandated by Aliyun RPC SignatureMethod
|
||||
).digest()
|
||||
return base64.b64encode(digest).decode("ascii")
|
||||
|
||||
@staticmethod
|
||||
def _percent_encode(value: str) -> str:
|
||||
return urllib.parse.quote(str(value), safe="~")
|
||||
|
||||
|
||||
class CaptchaService:
|
||||
def __init__(self, redis: Redis | None = None):
|
||||
self.redis = redis or get_redis()
|
||||
|
||||
async def create(self, identifier: str) -> bytes:
|
||||
code = "".join(secrets.choice(string.ascii_letters + string.digits) for _ in range(CAPTCHA_LENGTH))
|
||||
await self._set_cache(identifier, code)
|
||||
return self._render_gif(code)
|
||||
|
||||
async def validate(self, identifier: str | None, code: str | None, *, delete: bool) -> bool:
|
||||
if not code or not code.strip():
|
||||
return False
|
||||
key = self._captcha_key(identifier)
|
||||
cached = JavaRedisCodec.decode(await cast(Any, self.redis.get(key)))
|
||||
if cached is not None and delete:
|
||||
await cast(Any, self.redis.delete(key))
|
||||
return cached is not None and code.casefold() == str(cached).casefold()
|
||||
|
||||
async def set_sms_code(self, phone: str | None, code: str) -> None:
|
||||
await self._set_cache(f"sms:Validate:Code:{phone}", code)
|
||||
|
||||
async def validate_sms_code(self, phone: str | None, code: str | None, *, delete: bool = False) -> bool:
|
||||
return await self.validate(f"sms:Validate:Code:{phone}", code, delete=delete)
|
||||
|
||||
async def _set_cache(self, identifier: str, value: str) -> None:
|
||||
await cast(Any, self.redis.set)(
|
||||
self._captcha_key(identifier),
|
||||
JavaRedisCodec.encode(value),
|
||||
ex=CAPTCHA_TTL_SECONDS,
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _captcha_key(identifier: str | None) -> str:
|
||||
return f"sys:captcha:{'null' if identifier is None else identifier}"
|
||||
|
||||
@staticmethod
|
||||
def _render_gif(code: str) -> bytes:
|
||||
image = Image.new("RGB", (150, 40), (248, 248, 248))
|
||||
draw = ImageDraw.Draw(image)
|
||||
for _ in range(8):
|
||||
color = tuple(secrets.randbelow(150) for _ in range(3))
|
||||
draw.line(
|
||||
(
|
||||
secrets.randbelow(150),
|
||||
secrets.randbelow(40),
|
||||
secrets.randbelow(150),
|
||||
secrets.randbelow(40),
|
||||
),
|
||||
fill=color,
|
||||
width=1,
|
||||
)
|
||||
font = ImageFont.load_default(size=24)
|
||||
for index, character in enumerate(code):
|
||||
color = tuple(secrets.randbelow(120) for _ in range(3))
|
||||
draw.text((10 + index * 27, 6 + secrets.randbelow(5)), character, font=font, fill=color)
|
||||
output = io.BytesIO()
|
||||
image.save(output, format="GIF")
|
||||
return output.getvalue()
|
||||
|
||||
|
||||
class SecurityService:
|
||||
def __init__(
|
||||
self,
|
||||
repository: SecurityRepository,
|
||||
*,
|
||||
redis: Redis | None = None,
|
||||
captcha: CaptchaService | None = None,
|
||||
sms_sender: SmsSender | None = None,
|
||||
):
|
||||
self.repository = repository
|
||||
self.redis = redis or get_redis()
|
||||
self.captcha = captcha or CaptchaService(self.redis)
|
||||
self.sms_sender = sms_sender or AliyunSmsSender(repository, redis=self.redis)
|
||||
|
||||
async def login(self, dto: LoginRequest, request: Request) -> dict[str, Any]:
|
||||
password = await self._decrypt_and_validate_captcha(dto.password, dto.captcha_id)
|
||||
user = await self.repository.get_user_by_username(dto.username)
|
||||
if user is None or not bcrypt_matches(password, cast(str | None, user.get("password"))):
|
||||
raise AppError(ErrorCode.ACCOUNT_PASSWORD_ERROR)
|
||||
token = await self._create_token(int(user["id"]))
|
||||
await self.repository.session.commit()
|
||||
return {
|
||||
"token": token,
|
||||
"expire": TOKEN_EXPIRE_SECONDS,
|
||||
"clientHash": self._client_hash(request),
|
||||
}
|
||||
|
||||
async def register(self, dto: LoginRequest) -> None:
|
||||
if not await self.allow_user_register():
|
||||
raise AppError(10072)
|
||||
password = await self._decrypt_and_validate_captcha(dto.password, dto.captcha_id)
|
||||
if await self._mobile_registration_enabled():
|
||||
if dto.username is None or not PHONE_PATTERN.fullmatch(dto.username):
|
||||
raise AppError(10069)
|
||||
if not await self.captcha.validate_sms_code(dto.username, dto.mobile_captcha, delete=False):
|
||||
raise AppError(10075)
|
||||
if await self.repository.get_user_by_username(dto.username) is not None:
|
||||
raise AppError(10070)
|
||||
if not STRONG_PASSWORD.fullmatch(password):
|
||||
raise AppError(ErrorCode.PASSWORD_WEAK_ERROR)
|
||||
now = shanghai_now_naive()
|
||||
user_count = await self.repository.count_users()
|
||||
await self.repository.insert_user(
|
||||
user_id=snowflake.next_id(),
|
||||
username=dto.username,
|
||||
password=bcrypt_hash(password),
|
||||
super_admin=1 if user_count == 0 else 0,
|
||||
now=now,
|
||||
)
|
||||
await self.repository.session.commit()
|
||||
|
||||
async def change_password(
|
||||
self,
|
||||
user: AuthUser,
|
||||
dto: PasswordChangeRequest,
|
||||
accept_language: str | None = None,
|
||||
) -> None:
|
||||
self._require_not_blank(dto.password, "sysuser.password.require", accept_language)
|
||||
self._require_not_blank(dto.new_password, "sysuser.password.require", accept_language)
|
||||
assert dto.password is not None
|
||||
assert dto.new_password is not None
|
||||
row = await self.repository.get_user_by_id(user.id)
|
||||
if row is None:
|
||||
raise AppError(ErrorCode.TOKEN_INVALID)
|
||||
if not bcrypt_matches(dto.password, cast(str | None, row.get("password"))):
|
||||
raise AppError(10048)
|
||||
if not STRONG_PASSWORD.fullmatch(dto.new_password):
|
||||
raise AppError(ErrorCode.PASSWORD_WEAK_ERROR)
|
||||
now = shanghai_now_naive()
|
||||
await self.repository.update_password(
|
||||
user.id,
|
||||
bcrypt_hash(dto.new_password),
|
||||
now,
|
||||
preserve_audit_fields=True,
|
||||
)
|
||||
# SysUserService.changePassword commits before the non-transactional token service logs out.
|
||||
await self.repository.session.commit()
|
||||
await self.repository.expire_user_token(user.id, now - timedelta(minutes=1))
|
||||
await self.repository.session.commit()
|
||||
|
||||
async def retrieve_password(
|
||||
self,
|
||||
dto: RetrievePasswordRequest,
|
||||
accept_language: str | None = None,
|
||||
) -> None:
|
||||
if not await self._mobile_registration_enabled():
|
||||
raise AppError(10073)
|
||||
self._require_not_blank(dto.phone, "sysuser.password.require", accept_language)
|
||||
self._require_not_blank(dto.code, "sysuser.password.require", accept_language)
|
||||
self._require_not_blank(dto.password, "sysuser.password.require", accept_language)
|
||||
self._require_not_blank(dto.captcha_id, "sysuser.uuid.require", accept_language)
|
||||
assert dto.phone is not None
|
||||
assert dto.code is not None
|
||||
assert dto.password is not None
|
||||
assert dto.captcha_id is not None
|
||||
if not PHONE_PATTERN.fullmatch(dto.phone):
|
||||
raise AppError(10074)
|
||||
user = await self.repository.get_user_by_username(dto.phone)
|
||||
if user is None:
|
||||
raise AppError(10071)
|
||||
if not await self.captcha.validate_sms_code(dto.phone, dto.code, delete=False):
|
||||
raise AppError(10075)
|
||||
password = await self._decrypt_and_validate_captcha(dto.password, dto.captcha_id)
|
||||
if not STRONG_PASSWORD.fullmatch(password):
|
||||
raise AppError(ErrorCode.PASSWORD_WEAK_ERROR)
|
||||
await self.repository.update_password(int(user["id"]), bcrypt_hash(password), shanghai_now_naive())
|
||||
await self.repository.session.commit()
|
||||
|
||||
async def send_sms_verification(self, dto: SmsVerificationRequest) -> None:
|
||||
if not await self.captcha.validate(dto.captcha_id, dto.captcha, delete=False):
|
||||
raise AppError(10067)
|
||||
if not await self._mobile_registration_enabled():
|
||||
raise AppError(10068)
|
||||
phone_key = "null" if dto.phone is None else dto.phone
|
||||
last_send_key = f"sms:Validate:Code:{phone_key}:last_send_time"
|
||||
current_ms = int(time.time() * 1000)
|
||||
created = await cast(Any, self.redis.set)(last_send_key, str(current_ms), ex=60, nx=True)
|
||||
if not created:
|
||||
raw_last = await cast(Any, self.redis.get)(last_send_key)
|
||||
if raw_last is not None:
|
||||
last_ms = int(raw_last.decode() if isinstance(raw_last, bytes) else raw_last)
|
||||
difference = current_ms - last_ms
|
||||
if difference < 60_000:
|
||||
raise AppError(10060, params=(str(max(0, (60_000 - difference) // 1000)),))
|
||||
|
||||
today_key = f"sms:Validate:Code:{phone_key}:today_count"
|
||||
raw_count = await cast(Any, self.redis.get)(today_key)
|
||||
decoded_count = JavaRedisCodec.decode(raw_count)
|
||||
today_count = int(decoded_count or 0)
|
||||
raw_maximum = await self._get_param("server.sms_max_send_count", from_cache=True)
|
||||
maximum = int(raw_maximum) if raw_maximum is not None and raw_maximum != "" else 5
|
||||
if today_count >= maximum:
|
||||
raise AppError(10047)
|
||||
|
||||
code = "".join(secrets.choice(string.digits) for _ in range(6))
|
||||
await self.captcha.set_sms_code(dto.phone, code)
|
||||
new_count = await cast(Any, self.redis.incr)(today_key)
|
||||
if int(new_count) == 1:
|
||||
await cast(Any, self.redis.expire)(today_key, 24 * 60 * 60)
|
||||
try:
|
||||
await self.sms_sender.send_verification_code(dto.phone, code)
|
||||
except AppError:
|
||||
# Java raises connection-construction failures before entering its refundable send attempt.
|
||||
raise
|
||||
except Exception as exc:
|
||||
logger.warning("Aliyun SMS request failed", exc_info=exc)
|
||||
await cast(Any, self.redis.delete)(today_key)
|
||||
raise AppError(10055) from exc
|
||||
|
||||
async def public_config(self) -> dict[str, Any]:
|
||||
public_key = await self._get_param("server.public_key", from_cache=True)
|
||||
if public_key is None or not public_key.strip():
|
||||
raise AppError(10129)
|
||||
menu_config = await self._get_param("system-web.menu", from_cache=True)
|
||||
result: dict[str, Any] = {
|
||||
"enableMobileRegister": await self._mobile_registration_enabled(),
|
||||
"version": "0.9.5",
|
||||
"year": f"©{shanghai_now_naive().year}",
|
||||
"allowUserRegister": await self.allow_user_register(),
|
||||
"mobileAreaList": await self._dict_data_by_type("MOBILE_AREA"),
|
||||
"beianIcpNum": await self._get_param("server.beian_icp_num", from_cache=True),
|
||||
"beianGaNum": await self._get_param("server.beian_ga_num", from_cache=True),
|
||||
"name": await self._get_param("server.name", from_cache=True),
|
||||
"sm2PublicKey": public_key,
|
||||
}
|
||||
if menu_config is not None and menu_config.strip():
|
||||
result["systemWebMenu"] = json.loads(menu_config)
|
||||
return result
|
||||
|
||||
async def allow_user_register(self) -> bool:
|
||||
value = await self._get_param("server.allow_user_register", from_cache=True)
|
||||
if value == "true":
|
||||
return True
|
||||
return await self.repository.count_users() == 0
|
||||
|
||||
async def _create_token(self, user_id: int) -> str:
|
||||
now = shanghai_now_naive()
|
||||
expire_date = now + timedelta(seconds=TOKEN_EXPIRE_SECONDS)
|
||||
current = await self.repository.get_token_by_user_id(user_id, for_update=True)
|
||||
if current is None:
|
||||
token = generate_database_token()
|
||||
await self.repository.insert_token(
|
||||
token_id=snowflake.next_id(),
|
||||
user_id=user_id,
|
||||
token=token,
|
||||
now=now,
|
||||
expire_date=expire_date,
|
||||
)
|
||||
return token
|
||||
stored_expiry = self._datetime(current.get("expire_date"))
|
||||
token = str(current["token"])
|
||||
if stored_expiry is None or stored_expiry < now:
|
||||
token = generate_database_token()
|
||||
await self.repository.update_token(
|
||||
token_id=int(current["id"]),
|
||||
token=token,
|
||||
now=now,
|
||||
expire_date=expire_date,
|
||||
)
|
||||
return token
|
||||
|
||||
async def _decrypt_and_validate_captcha(
|
||||
self,
|
||||
encrypted_password: str | None,
|
||||
captcha_id: str | None,
|
||||
) -> str:
|
||||
private_key = await self._get_param("server.private_key", from_cache=True)
|
||||
if private_key is None or not private_key.strip():
|
||||
raise AppError(10129)
|
||||
try:
|
||||
if encrypted_password is None:
|
||||
raise ValueError("encrypted password is null")
|
||||
content = sm2_decrypt_c1c3c2(private_key, encrypted_password)
|
||||
except Exception as exc:
|
||||
raise AppError(10130) from exc
|
||||
if len(content) > CAPTCHA_LENGTH:
|
||||
embedded_captcha = content[:CAPTCHA_LENGTH]
|
||||
if not await self.captcha.validate(captcha_id, embedded_captcha, delete=True):
|
||||
raise AppError(10067)
|
||||
return content[CAPTCHA_LENGTH:]
|
||||
if content:
|
||||
raise AppError(10067)
|
||||
raise AppError(10130)
|
||||
|
||||
async def _mobile_registration_enabled(self) -> bool:
|
||||
value = await self._get_param("server.enable_mobile_register", from_cache=True)
|
||||
if value is None or not value.strip():
|
||||
return False
|
||||
try:
|
||||
parsed = json.loads(value.lower())
|
||||
except json.JSONDecodeError as exc:
|
||||
raise AppError(ErrorCode.PARAMS_GET_ERROR) from exc
|
||||
return bool(parsed)
|
||||
|
||||
async def _get_param(self, code: str, *, from_cache: bool) -> str | None:
|
||||
if from_cache:
|
||||
cached = JavaRedisCodec.decode(await cast(Any, self.redis.hget)("sys:params", code))
|
||||
if cached is not None:
|
||||
return str(cached)
|
||||
value = await self.repository.get_param_value(code)
|
||||
if from_cache and value is not None:
|
||||
await cast(Any, self.redis.hset)("sys:params", code, JavaRedisCodec.encode(value))
|
||||
await cast(Any, self.redis.expire)("sys:params", 24 * 60 * 60)
|
||||
return value
|
||||
|
||||
async def _dict_data_by_type(self, dict_type: str) -> list[dict[str, Any]]:
|
||||
key = f"sys:dict:data:{dict_type}"
|
||||
cached = JavaRedisCodec.decode(await cast(Any, self.redis.get)(key))
|
||||
if isinstance(cached, list):
|
||||
return cast(list[dict[str, Any]], cached)
|
||||
values = await self.repository.get_mobile_area_items()
|
||||
await cast(Any, self.redis.set)(
|
||||
key,
|
||||
JavaRedisCodec.encode(
|
||||
values,
|
||||
item_java_type="xiaozhi.modules.sys.vo.SysDictDataItem",
|
||||
),
|
||||
ex=24 * 60 * 60,
|
||||
)
|
||||
return values
|
||||
|
||||
@staticmethod
|
||||
def _client_hash(request: Request) -> str:
|
||||
user_agent = request.headers.get("User-Agent", "").lower()
|
||||
forwarded_headers = (
|
||||
"x-forwarded-for",
|
||||
"Proxy-Client-IP",
|
||||
"WL-Proxy-Client-IP",
|
||||
"HTTP_CLIENT_IP",
|
||||
"HTTP_X_FORWARDED_FOR",
|
||||
)
|
||||
ip_address = next(
|
||||
(
|
||||
value
|
||||
for header in forwarded_headers
|
||||
if (value := request.headers.get(header)) and value.casefold() != "unknown"
|
||||
),
|
||||
request.client.host if request.client else "",
|
||||
)
|
||||
date = shanghai_now_naive().strftime("%Y-%m-%d")
|
||||
return hashlib.md5( # noqa: S324 - Java clientHash compatibility requires MD5
|
||||
f"{ip_address}{date}{user_agent}".encode(), usedforsecurity=False
|
||||
).hexdigest()
|
||||
|
||||
@staticmethod
|
||||
def _datetime(value: Any) -> datetime | None:
|
||||
if value is None or isinstance(value, datetime):
|
||||
return value
|
||||
if isinstance(value, str):
|
||||
return datetime.fromisoformat(value)
|
||||
raise TypeError(f"Unsupported database datetime value: {type(value).__name__}")
|
||||
|
||||
@staticmethod
|
||||
def _require_not_blank(value: str | None, key: str, accept_language: str | None) -> None:
|
||||
if value is None or not value.strip():
|
||||
raise AppError(500, validation_message(key, accept_language))
|
||||
@@ -0,0 +1,715 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import base64
|
||||
import hashlib
|
||||
import hmac
|
||||
import json
|
||||
import logging
|
||||
import re
|
||||
import secrets
|
||||
import string
|
||||
import time
|
||||
import uuid
|
||||
from datetime import datetime
|
||||
from typing import Any, cast
|
||||
from zoneinfo import ZoneInfo
|
||||
|
||||
import httpx
|
||||
from redis.asyncio import Redis
|
||||
from websockets.asyncio.client import connect
|
||||
|
||||
from app.core.config import get_settings
|
||||
from app.core.crypto import bcrypt_hash
|
||||
from app.core.errors import AppError, ErrorCode
|
||||
from app.core.ids import snowflake
|
||||
from app.core.redis import JavaRedisCodec, get_redis
|
||||
from app.core.security import AuthUser, shanghai_now_naive
|
||||
from app.repositories.sys import SysRepository
|
||||
from app.schemas.sys import DictDataPayload, DictTypePayload, EmitServerActionRequest, SysParamPayload
|
||||
from app.services.java_validation import validation_message
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
WS_PATTERN = re.compile(r"^wss?://[\w.-]+(?:\.[\w.-]+)*(?::\d+)?(?:/[\w.-]*)*$")
|
||||
|
||||
|
||||
class AdminService:
|
||||
def __init__(self, repository: SysRepository):
|
||||
self.repository = repository
|
||||
|
||||
async def page_users(self, *, mobile: str | None, page: int, limit: int) -> dict[str, Any]:
|
||||
rows, total = await self.repository.page_users(
|
||||
mobile=mobile,
|
||||
page=max(1, page),
|
||||
limit=max(0, limit),
|
||||
)
|
||||
values = [
|
||||
{
|
||||
"deviceCount": str(row.get("device_count") or 0),
|
||||
"mobile": row.get("username"),
|
||||
"status": row.get("status"),
|
||||
"userid": str(row["id"]),
|
||||
"createDate": row.get("create_date"),
|
||||
}
|
||||
for row in rows
|
||||
]
|
||||
return {"list": values, "total": total}
|
||||
|
||||
async def reset_password(self, user_id: int, user: AuthUser) -> str:
|
||||
password = self._generate_password()
|
||||
await self.repository.reset_user_password(user_id, bcrypt_hash(password), user.id, shanghai_now_naive())
|
||||
await self.repository.session.commit()
|
||||
return password
|
||||
|
||||
async def delete_user(self, user_id: int) -> None:
|
||||
try:
|
||||
await self.repository.delete_user_cascade(user_id)
|
||||
await self.repository.session.commit()
|
||||
except Exception:
|
||||
await self.repository.session.rollback()
|
||||
raise
|
||||
|
||||
async def change_status(self, status: int, user_ids: list[str], user: AuthUser) -> None:
|
||||
# SysUserServiceImpl.changeStatus has an outer Spring transaction: a later
|
||||
# parse/update failure rolls back every earlier item in the same request.
|
||||
try:
|
||||
for value in user_ids:
|
||||
await self.repository.change_user_status(status, [int(value)], user.id, shanghai_now_naive())
|
||||
await self.repository.session.commit()
|
||||
except Exception:
|
||||
await self.repository.session.rollback()
|
||||
raise
|
||||
|
||||
async def page_devices(self, *, keywords: str | None, page: int, limit: int) -> dict[str, Any]:
|
||||
rows, total = await self.repository.page_devices(
|
||||
keywords=keywords,
|
||||
page=max(1, page),
|
||||
limit=max(0, limit),
|
||||
)
|
||||
result = []
|
||||
for row in rows:
|
||||
result.append(
|
||||
{
|
||||
"appVersion": row.get("app_version"),
|
||||
"bindUserName": row.get("bind_user_name"),
|
||||
"deviceType": row.get("board"),
|
||||
"board": row.get("board"),
|
||||
"id": row.get("id"),
|
||||
"macAddress": row.get("mac_address"),
|
||||
"alias": row.get("alias"),
|
||||
"otaUpgrade": None,
|
||||
"recentChatTime": self._short_time(cast(datetime | str | None, row.get("update_date"))),
|
||||
"lastConnectedAtTimestamp": self._timestamp_ms(
|
||||
cast(datetime | str | None, row.get("last_connected_at"))
|
||||
),
|
||||
"createDateTimestamp": self._timestamp_ms(
|
||||
cast(datetime | str | None, row.get("create_date"))
|
||||
),
|
||||
"createDate": self._utc_datetime_string(
|
||||
cast(datetime | str | None, row.get("create_date"))
|
||||
),
|
||||
}
|
||||
)
|
||||
return {"list": result, "total": total}
|
||||
|
||||
@staticmethod
|
||||
def _generate_password() -> str:
|
||||
characters = string.ascii_letters + string.digits + "!@#$%^&*()"
|
||||
values = [
|
||||
secrets.choice(string.digits),
|
||||
secrets.choice(string.ascii_lowercase),
|
||||
secrets.choice(string.ascii_uppercase),
|
||||
secrets.choice("!@#$%^&*()"),
|
||||
]
|
||||
values.extend(secrets.choice(characters) for _ in range(8))
|
||||
secrets.SystemRandom().shuffle(values)
|
||||
return "".join(values)
|
||||
|
||||
@staticmethod
|
||||
def _timestamp_ms(value: datetime | str | None) -> int | None:
|
||||
value = AdminService._database_datetime(value)
|
||||
if value is None:
|
||||
return None
|
||||
timezone = ZoneInfo(get_settings().timezone)
|
||||
localized = value if value.tzinfo else value.replace(tzinfo=timezone)
|
||||
return int(localized.timestamp() * 1000)
|
||||
|
||||
@staticmethod
|
||||
def _short_time(value: datetime | str | None) -> str | None:
|
||||
value = AdminService._database_datetime(value)
|
||||
if value is None:
|
||||
return None
|
||||
now = shanghai_now_naive()
|
||||
if value.tzinfo:
|
||||
value = value.astimezone(ZoneInfo(get_settings().timezone)).replace(tzinfo=None)
|
||||
seconds = int((now - value).total_seconds())
|
||||
if seconds <= 10:
|
||||
return "刚刚"
|
||||
if seconds < 60:
|
||||
return f"{seconds}秒前"
|
||||
if seconds < 3600:
|
||||
return f"{seconds // 60}分钟前"
|
||||
if seconds < 86400:
|
||||
return f"{seconds // 3600}小时前"
|
||||
if seconds < 604800:
|
||||
return f"{seconds // 86400}天前"
|
||||
return value.strftime("%Y-%m-%d %H:%M:%S")
|
||||
|
||||
@staticmethod
|
||||
def _utc_datetime_string(value: datetime | str | None) -> str | None:
|
||||
value = AdminService._database_datetime(value)
|
||||
if value is None:
|
||||
return None
|
||||
timezone = ZoneInfo(get_settings().timezone)
|
||||
localized = value if value.tzinfo else value.replace(tzinfo=timezone)
|
||||
return localized.astimezone(ZoneInfo("UTC")).strftime("%Y-%m-%d %H:%M:%S")
|
||||
|
||||
@staticmethod
|
||||
def _database_datetime(value: datetime | str | None) -> datetime | None:
|
||||
if isinstance(value, str):
|
||||
return datetime.fromisoformat(value)
|
||||
return value
|
||||
|
||||
|
||||
class ParamExternalValidator:
|
||||
def __init__(self, client: httpx.AsyncClient | None = None):
|
||||
self.client = client
|
||||
|
||||
async def validate(self, code: str, value: str) -> None:
|
||||
if code == "server.websocket":
|
||||
await self._websockets(value)
|
||||
elif code == "server.ota":
|
||||
await self._http_endpoint(value, kind="ota")
|
||||
elif code == "server.mcp_endpoint":
|
||||
await self._http_endpoint(value, kind="mcp")
|
||||
elif code == "server.voice_print":
|
||||
await self._http_endpoint(value, kind="voiceprint")
|
||||
elif code == "server.mqtt_signature_key":
|
||||
self._mqtt_secret(value)
|
||||
|
||||
async def _websockets(self, value: str) -> None:
|
||||
urls = value.split(";")
|
||||
while urls and urls[-1] == "":
|
||||
urls.pop()
|
||||
if not urls:
|
||||
raise AppError(10098)
|
||||
for raw_url in urls:
|
||||
if not raw_url.strip():
|
||||
continue
|
||||
if "localhost" in raw_url or "127.0.0.1" in raw_url:
|
||||
raise AppError(10099)
|
||||
if not WS_PATTERN.fullmatch(raw_url.strip()):
|
||||
raise AppError(10100)
|
||||
try:
|
||||
async with connect(raw_url, open_timeout=5):
|
||||
pass
|
||||
except Exception as exc:
|
||||
raise AppError(10101) from exc
|
||||
|
||||
async def _http_endpoint(self, value: str, *, kind: str) -> None:
|
||||
if not value.strip() or value == "null":
|
||||
return
|
||||
if "localhost" in value or "127.0.0.1" in value:
|
||||
raise AppError({"ota": 10103, "mcp": 10110, "voiceprint": 10116}[kind])
|
||||
if kind == "ota":
|
||||
if not value.lower().startswith("http"):
|
||||
raise AppError(10104)
|
||||
if not value.endswith("/ota/"):
|
||||
raise AppError(10105)
|
||||
elif kind == "mcp":
|
||||
if "key" not in value.lower():
|
||||
raise AppError(10111)
|
||||
else:
|
||||
if "key" not in value.lower():
|
||||
raise AppError(10117)
|
||||
if not value.lower().startswith("http"):
|
||||
raise AppError(10118)
|
||||
|
||||
final_code = {"ota": 10108, "mcp": 10114, "voiceprint": 10121}[kind]
|
||||
marker = {"ota": "OTA", "mcp": "success", "voiceprint": "healthy"}[kind]
|
||||
try:
|
||||
if self.client is not None:
|
||||
response = await self.client.get(value)
|
||||
else:
|
||||
async with httpx.AsyncClient(timeout=get_settings().external_request_timeout_seconds) as client:
|
||||
response = await client.get(value)
|
||||
if response.status_code != 200 or marker not in response.text:
|
||||
raise ValueError("external endpoint response did not match Java validation")
|
||||
except Exception as exc:
|
||||
raise AppError(final_code) from exc
|
||||
|
||||
@staticmethod
|
||||
def _mqtt_secret(secret: str) -> None:
|
||||
if not secret.strip() or secret == "null": # noqa: S105 - sentinel value from the Java parameter table
|
||||
raise AppError(10122)
|
||||
if len(secret) < 8:
|
||||
raise AppError(10123)
|
||||
if not re.search(r"[a-z]", secret) or not re.search(r"[A-Z]", secret):
|
||||
raise AppError(10124)
|
||||
lowered = secret.lower()
|
||||
if any(weak in lowered for weak in ("test", "1234", "admin", "password", "qwerty", "xiaozhi")):
|
||||
raise AppError(10125)
|
||||
|
||||
|
||||
class SysParamService:
|
||||
def __init__(
|
||||
self,
|
||||
repository: SysRepository,
|
||||
*,
|
||||
redis: Redis | None = None,
|
||||
validator: ParamExternalValidator | None = None,
|
||||
):
|
||||
self.repository = repository
|
||||
self.redis = redis or get_redis()
|
||||
self.validator = validator or ParamExternalValidator()
|
||||
|
||||
async def page(
|
||||
self,
|
||||
*,
|
||||
param_code: str | None,
|
||||
page: int,
|
||||
limit: int,
|
||||
order_field: str | None,
|
||||
order: str | None,
|
||||
) -> dict[str, Any]:
|
||||
rows, total = await self.repository.page_params(
|
||||
param_code=param_code,
|
||||
page=max(1, page),
|
||||
limit=max(0, limit),
|
||||
order_field=order_field,
|
||||
order=order,
|
||||
)
|
||||
return {"list": [self._param_dto(row) for row in rows], "total": total}
|
||||
|
||||
async def get(self, param_id: int) -> dict[str, Any] | None:
|
||||
row = await self.repository.get_param(param_id)
|
||||
return None if row is None else self._param_dto(row)
|
||||
|
||||
async def save(
|
||||
self,
|
||||
dto: SysParamPayload,
|
||||
user: AuthUser,
|
||||
accept_language: str | None = None,
|
||||
) -> None:
|
||||
self._validate_group(dto, update=False, accept_language=accept_language)
|
||||
self._validate_value(dto)
|
||||
assert dto.param_code is not None
|
||||
assert dto.param_value is not None
|
||||
assert dto.value_type is not None
|
||||
await self.repository.insert_param(
|
||||
param_id=snowflake.next_id(),
|
||||
param_code=dto.param_code,
|
||||
param_value=dto.param_value,
|
||||
value_type=dto.value_type,
|
||||
remark=dto.remark,
|
||||
user_id=user.id,
|
||||
now=shanghai_now_naive(),
|
||||
)
|
||||
await self._cache_set(dto.param_code, dto.param_value)
|
||||
await self.repository.session.commit()
|
||||
|
||||
async def update(
|
||||
self,
|
||||
dto: SysParamPayload,
|
||||
user: AuthUser,
|
||||
accept_language: str | None = None,
|
||||
) -> None:
|
||||
self._validate_group(dto, update=True, accept_language=accept_language)
|
||||
assert dto.id is not None
|
||||
assert dto.param_code is not None
|
||||
assert dto.param_value is not None
|
||||
assert dto.value_type is not None
|
||||
# These checks live in the Java controller and therefore run before
|
||||
# SysParamsService.update validates the declared value type.
|
||||
await self.validator.validate(dto.param_code, dto.param_value)
|
||||
if dto.param_code == "system-web.menu":
|
||||
await self._update_system_web_menu(dto.param_value, user)
|
||||
else:
|
||||
self._validate_value(dto)
|
||||
await self.repository.update_param(
|
||||
param_id=dto.id,
|
||||
param_code=dto.param_code,
|
||||
param_value=dto.param_value,
|
||||
value_type=dto.value_type,
|
||||
remark=dto.remark,
|
||||
user_id=user.id,
|
||||
now=shanghai_now_naive(),
|
||||
)
|
||||
await self._cache_set(dto.param_code, dto.param_value)
|
||||
await self.repository.session.commit()
|
||||
|
||||
async def delete(self, ids: list[str]) -> None:
|
||||
if not ids:
|
||||
raise AppError(10001, "id")
|
||||
parsed_ids = [int(value) for value in ids]
|
||||
codes = await self.repository.param_codes_for_ids(parsed_ids)
|
||||
if codes:
|
||||
await cast(Any, self.redis.hdel)("sys:params", *codes)
|
||||
await self.repository.delete_params(parsed_ids)
|
||||
await self.repository.session.commit()
|
||||
|
||||
async def get_value(self, code: str, *, from_cache: bool = True) -> str | None:
|
||||
if from_cache:
|
||||
cached = JavaRedisCodec.decode(await cast(Any, self.redis.hget)("sys:params", code))
|
||||
if cached is not None:
|
||||
return str(cached)
|
||||
value = await self.repository.get_param_value(code)
|
||||
if value is not None and from_cache:
|
||||
await self._cache_set(code, value)
|
||||
return value
|
||||
|
||||
async def config_rows(self) -> list[dict[str, Any]]:
|
||||
return await self.repository.list_config_params()
|
||||
|
||||
async def _update_system_web_menu(self, config_json: str, user: AuthUser) -> None:
|
||||
current_config = await self.repository.get_param_value("system-web.menu")
|
||||
try:
|
||||
current = json.loads(current_config) if current_config and current_config.strip() else None
|
||||
updated = json.loads(config_json) if config_json.strip() else None
|
||||
except json.JSONDecodeError as exc:
|
||||
raise AppError(ErrorCode.PARAM_JSON_INVALID) from exc
|
||||
if isinstance(current, dict) and isinstance(updated, dict):
|
||||
current_features = current.get("features")
|
||||
updated_features = updated.get("features")
|
||||
# Java only evaluates addressBook when both feature maps are present.
|
||||
if isinstance(current_features, dict) and isinstance(updated_features, dict):
|
||||
current_address = current_features.get("addressBook")
|
||||
updated_address = updated_features.get("addressBook")
|
||||
current_enabled = self._java_enabled(current_address)
|
||||
updated_enabled = self._java_enabled(updated_address)
|
||||
if current_enabled and not updated_enabled:
|
||||
await self.repository.delete_plugin_mapping_by_plugin_id("SYSTEM_PLUGIN_CALL_DEVICE")
|
||||
await self.repository.update_param_value_by_code(
|
||||
"system-web.menu", config_json, user.id, shanghai_now_naive()
|
||||
)
|
||||
await self._cache_set("system-web.menu", config_json)
|
||||
|
||||
async def _cache_set(self, code: str, value: str) -> None:
|
||||
await cast(Any, self.redis.hset)("sys:params", code, JavaRedisCodec.encode(value))
|
||||
await cast(Any, self.redis.expire)("sys:params", 24 * 60 * 60)
|
||||
|
||||
@staticmethod
|
||||
def _java_enabled(address_book: Any) -> bool:
|
||||
if not isinstance(address_book, dict):
|
||||
return False
|
||||
value = address_book.get("enabled")
|
||||
if value is None:
|
||||
return False
|
||||
if not isinstance(value, bool):
|
||||
# The Java implementation casts the JSON value to Boolean.
|
||||
raise TypeError("addressBook.enabled must be a boolean")
|
||||
return value
|
||||
|
||||
@staticmethod
|
||||
def _validate_value(dto: SysParamPayload) -> None:
|
||||
assert dto.param_value is not None
|
||||
assert dto.value_type is not None
|
||||
if not dto.param_value.strip():
|
||||
raise AppError(ErrorCode.PARAM_VALUE_NULL)
|
||||
if not dto.value_type.strip():
|
||||
raise AppError(ErrorCode.PARAM_TYPE_NULL)
|
||||
value_type = dto.value_type.lower()
|
||||
if value_type in {"string", "array"}:
|
||||
return
|
||||
if value_type == "number":
|
||||
try:
|
||||
float(dto.param_value)
|
||||
except ValueError as exc:
|
||||
raise AppError(ErrorCode.PARAM_NUMBER_INVALID) from exc
|
||||
return
|
||||
if value_type == "boolean":
|
||||
if dto.param_value.lower() not in {"true", "false"}:
|
||||
raise AppError(ErrorCode.PARAM_BOOLEAN_INVALID)
|
||||
return
|
||||
if value_type == "json":
|
||||
stripped = dto.param_value.strip()
|
||||
if not stripped.startswith("{") or not stripped.endswith("}"):
|
||||
raise AppError(ErrorCode.PARAM_JSON_INVALID)
|
||||
try:
|
||||
json.loads(dto.param_value)
|
||||
except json.JSONDecodeError as exc:
|
||||
raise AppError(ErrorCode.PARAM_JSON_INVALID) from exc
|
||||
return
|
||||
raise AppError(ErrorCode.PARAM_TYPE_INVALID)
|
||||
|
||||
@staticmethod
|
||||
def _validate_group(
|
||||
dto: SysParamPayload,
|
||||
*,
|
||||
update: bool,
|
||||
accept_language: str | None,
|
||||
) -> None:
|
||||
def fail(key: str) -> None:
|
||||
raise AppError(500, validation_message(key, accept_language))
|
||||
|
||||
if update and dto.id is None:
|
||||
fail("id.require")
|
||||
if not update and dto.id is not None:
|
||||
fail("id.null")
|
||||
if dto.param_code is None or not dto.param_code.strip():
|
||||
fail("sysparams.paramcode.require")
|
||||
if dto.param_value is None or not dto.param_value.strip():
|
||||
fail("sysparams.paramvalue.require")
|
||||
if dto.value_type is None or not dto.value_type.strip():
|
||||
fail("sysparams.valuetype.require")
|
||||
if dto.value_type not in {"string", "number", "boolean", "array", "json"}:
|
||||
fail("sysparams.valuetype.pattern")
|
||||
|
||||
@staticmethod
|
||||
def _param_dto(row: dict[str, Any]) -> dict[str, Any]:
|
||||
return {
|
||||
"id": row.get("id"),
|
||||
"paramCode": row.get("param_code"),
|
||||
"paramValue": row.get("param_value"),
|
||||
"valueType": row.get("value_type"),
|
||||
"remark": row.get("remark"),
|
||||
"createDate": row.get("create_date"),
|
||||
"updateDate": row.get("update_date"),
|
||||
}
|
||||
|
||||
|
||||
class DictService:
|
||||
def __init__(self, repository: SysRepository, *, redis: Redis | None = None):
|
||||
self.repository = repository
|
||||
self.redis = redis or get_redis()
|
||||
|
||||
async def page_types(
|
||||
self,
|
||||
*,
|
||||
dict_type: str | None,
|
||||
dict_name: str | None,
|
||||
page: int,
|
||||
limit: int,
|
||||
) -> dict[str, Any]:
|
||||
rows, total = await self.repository.page_dict_types(
|
||||
dict_type=dict_type,
|
||||
dict_name=dict_name,
|
||||
page=max(1, page),
|
||||
limit=max(0, limit),
|
||||
)
|
||||
return {"list": [self._type_vo(row, include_names=True) for row in rows], "total": total}
|
||||
|
||||
async def get_type(self, type_id: int) -> dict[str, Any]:
|
||||
row = await self.repository.get_dict_type(type_id)
|
||||
if row is None:
|
||||
raise AppError(10076)
|
||||
return self._type_vo(row, include_names=False)
|
||||
|
||||
async def save_type(self, dto: DictTypePayload, user: AuthUser) -> None:
|
||||
if await self.repository.dict_type_exists(dto.dict_type):
|
||||
raise AppError(10077)
|
||||
await self.repository.insert_dict_type(
|
||||
type_id=dto.id if dto.id is not None else snowflake.next_id(),
|
||||
dict_type=dto.dict_type,
|
||||
dict_name=dto.dict_name,
|
||||
remark=dto.remark,
|
||||
sort=dto.sort,
|
||||
user_id=user.id,
|
||||
now=shanghai_now_naive(),
|
||||
)
|
||||
await self.repository.session.commit()
|
||||
|
||||
async def update_type(self, dto: DictTypePayload, user: AuthUser) -> None:
|
||||
if await self.repository.dict_type_exists(dto.dict_type, exclude_id=dto.id):
|
||||
raise AppError(10077)
|
||||
await self.repository.update_dict_type(
|
||||
type_id=dto.id,
|
||||
dict_type=dto.dict_type,
|
||||
dict_name=dto.dict_name,
|
||||
remark=dto.remark,
|
||||
sort=dto.sort,
|
||||
user_id=user.id,
|
||||
now=shanghai_now_naive(),
|
||||
)
|
||||
await self.repository.session.commit()
|
||||
|
||||
async def delete_types(self, ids: list[int]) -> None:
|
||||
await self.repository.delete_dict_types(ids)
|
||||
await self.repository.session.commit()
|
||||
|
||||
async def page_data(
|
||||
self,
|
||||
*,
|
||||
dict_type_id: int,
|
||||
dict_label: str | None,
|
||||
dict_value: str | None,
|
||||
page: int,
|
||||
limit: int,
|
||||
) -> dict[str, Any]:
|
||||
rows, total = await self.repository.page_dict_data(
|
||||
dict_type_id=dict_type_id,
|
||||
dict_label=dict_label,
|
||||
dict_value=dict_value,
|
||||
page=max(1, page),
|
||||
limit=max(0, limit),
|
||||
)
|
||||
return {"list": [self._data_vo(row, include_names=True) for row in rows], "total": total}
|
||||
|
||||
async def get_data(self, data_id: int) -> dict[str, Any] | None:
|
||||
row = await self.repository.get_dict_data(data_id)
|
||||
return None if row is None else self._data_vo(row, include_names=False)
|
||||
|
||||
async def save_data(self, dto: DictDataPayload, user: AuthUser) -> None:
|
||||
# Java compares dict_label against dictValue here; retain that behavior for compatibility.
|
||||
if await self.repository.dict_data_label_exists(dto.dict_type_id, dto.dict_value):
|
||||
raise AppError(10128)
|
||||
await self.repository.insert_dict_data(
|
||||
data_id=dto.id if dto.id is not None else snowflake.next_id(),
|
||||
dict_type_id=dto.dict_type_id,
|
||||
dict_label=dto.dict_label,
|
||||
dict_value=dto.dict_value,
|
||||
remark=dto.remark,
|
||||
sort=dto.sort,
|
||||
user_id=user.id,
|
||||
now=shanghai_now_naive(),
|
||||
)
|
||||
await self._clear_dict_cache(dto.dict_type_id)
|
||||
await self.repository.session.commit()
|
||||
|
||||
async def update_data(self, dto: DictDataPayload, user: AuthUser) -> None:
|
||||
if await self.repository.dict_data_label_exists(dto.dict_type_id, dto.dict_value, exclude_id=dto.id):
|
||||
raise AppError(10128)
|
||||
await self.repository.update_dict_data(
|
||||
data_id=dto.id,
|
||||
dict_type_id=dto.dict_type_id,
|
||||
dict_label=dto.dict_label,
|
||||
dict_value=dto.dict_value,
|
||||
remark=dto.remark,
|
||||
sort=dto.sort,
|
||||
user_id=user.id,
|
||||
now=shanghai_now_naive(),
|
||||
)
|
||||
await self._clear_dict_cache(dto.dict_type_id)
|
||||
await self.repository.session.commit()
|
||||
|
||||
async def delete_data(self, ids: list[int]) -> None:
|
||||
if ids:
|
||||
codes = await self.repository.dict_type_codes_for_data_ids(ids)
|
||||
if codes:
|
||||
await cast(Any, self.redis.delete)(*[f"sys:dict:data:{code}" for code in codes])
|
||||
await self.repository.delete_dict_data(ids)
|
||||
await self.repository.session.commit()
|
||||
|
||||
async def items(self, dict_type: str) -> list[dict[str, Any]] | None:
|
||||
if not dict_type.strip():
|
||||
return None
|
||||
key = f"sys:dict:data:{dict_type}"
|
||||
cached = JavaRedisCodec.decode(await cast(Any, self.redis.get)(key))
|
||||
if isinstance(cached, list):
|
||||
return cast(list[dict[str, Any]], cached)
|
||||
rows = await self.repository.dict_items(dict_type)
|
||||
await cast(Any, self.redis.set)(
|
||||
key,
|
||||
JavaRedisCodec.encode(
|
||||
rows,
|
||||
item_java_type="xiaozhi.modules.sys.vo.SysDictDataItem",
|
||||
),
|
||||
ex=24 * 60 * 60,
|
||||
)
|
||||
return rows
|
||||
|
||||
async def _clear_dict_cache(self, type_id: int | None) -> None:
|
||||
dict_type = await self.repository.dict_type_code(type_id)
|
||||
if dict_type is not None:
|
||||
await cast(Any, self.redis.delete)(f"sys:dict:data:{dict_type}")
|
||||
|
||||
@staticmethod
|
||||
def _type_vo(row: dict[str, Any], *, include_names: bool) -> dict[str, Any]:
|
||||
return {
|
||||
"id": row.get("id"),
|
||||
"dictType": row.get("dict_type"),
|
||||
"dictName": row.get("dict_name"),
|
||||
"remark": row.get("remark"),
|
||||
"sort": row.get("sort"),
|
||||
"creator": row.get("creator"),
|
||||
"creatorName": row.get("creator_name") if include_names else None,
|
||||
"createDate": row.get("create_date"),
|
||||
"updater": row.get("updater"),
|
||||
"updaterName": row.get("updater_name") if include_names else None,
|
||||
"updateDate": row.get("update_date"),
|
||||
}
|
||||
|
||||
@staticmethod
|
||||
def _data_vo(row: dict[str, Any], *, include_names: bool) -> dict[str, Any]:
|
||||
return {
|
||||
"id": row.get("id"),
|
||||
"dictTypeId": row.get("dict_type_id"),
|
||||
"dictLabel": row.get("dict_label"),
|
||||
"dictValue": row.get("dict_value"),
|
||||
"remark": row.get("remark"),
|
||||
"sort": row.get("sort"),
|
||||
"creator": row.get("creator"),
|
||||
"creatorName": row.get("creator_name") if include_names else None,
|
||||
"createDate": row.get("create_date"),
|
||||
"updater": row.get("updater"),
|
||||
"updaterName": row.get("updater_name") if include_names else None,
|
||||
"updateDate": row.get("update_date"),
|
||||
}
|
||||
|
||||
|
||||
class ServerActionService:
|
||||
def __init__(self, param_service: SysParamService, *, redis: Redis | None = None):
|
||||
self.param_service = param_service
|
||||
self.redis = redis or get_redis()
|
||||
|
||||
async def server_list(self) -> list[str]:
|
||||
value = await self.param_service.get_value("server.websocket", from_cache=True)
|
||||
if value is None or not value.strip():
|
||||
return []
|
||||
values = value.split(";")
|
||||
while values and values[-1] == "":
|
||||
values.pop()
|
||||
return values
|
||||
|
||||
async def emit(self, dto: EmitServerActionRequest) -> bool:
|
||||
action = dto.action.lower() if dto.action is not None else None
|
||||
if action not in {"restart", "update_config"}:
|
||||
raise AppError(10095)
|
||||
websocket_text = await self.param_service.get_value("server.websocket", from_cache=True)
|
||||
if websocket_text is None or not websocket_text.strip():
|
||||
raise AppError(10096)
|
||||
if dto.target_ws not in websocket_text.split(";"):
|
||||
raise AppError(10097)
|
||||
payload_secret = await self.param_service.get_value("server.secret", from_cache=True)
|
||||
device_id = str(uuid.uuid4())
|
||||
client_id = str(uuid.uuid4())
|
||||
await cast(Any, self.redis.set)(
|
||||
f"tmp_register_mac:{device_id}",
|
||||
JavaRedisCodec.encode("true"),
|
||||
ex=300,
|
||||
)
|
||||
authentication_secret = await self.param_service.get_value("server.secret", from_cache=False)
|
||||
if authentication_secret is None or not authentication_secret.strip():
|
||||
raise AppError(10045)
|
||||
timestamp = int(time.time())
|
||||
content = f"{client_id}|{device_id}|{timestamp}"
|
||||
signature = hmac.new(authentication_secret.encode(), content.encode(), digestmod=hashlib.sha256).digest()
|
||||
token = base64.urlsafe_b64encode(signature).rstrip(b"=").decode() + f".{timestamp}"
|
||||
headers = {
|
||||
"device-id": device_id,
|
||||
"client-id": client_id,
|
||||
"authorization": f"Bearer {token}",
|
||||
}
|
||||
if payload_secret is None:
|
||||
raise AppError(10045)
|
||||
payload = {"type": "server", "action": action, "content": {"secret": payload_secret}}
|
||||
try:
|
||||
async with connect(dto.target_ws, additional_headers=headers, open_timeout=3) as websocket:
|
||||
await websocket.send(json.dumps(payload, ensure_ascii=False, separators=(",", ":")))
|
||||
deadline = time.monotonic() + 120
|
||||
while True:
|
||||
remaining = deadline - time.monotonic()
|
||||
if remaining <= 0:
|
||||
raise TimeoutError
|
||||
raw = await asyncio.wait_for(websocket.recv(), timeout=remaining)
|
||||
response = json.loads(raw)
|
||||
if (
|
||||
isinstance(response, dict)
|
||||
and response.get("status") == "success"
|
||||
and response.get("type") == "server"
|
||||
and isinstance(response.get("content"), dict)
|
||||
and response["content"].get("action") is not None
|
||||
):
|
||||
return True
|
||||
except Exception as exc:
|
||||
raise AppError(10045) from exc
|
||||
@@ -0,0 +1,50 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
|
||||
from sqlalchemy import text
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
|
||||
from app.core.redis import java_hget, java_hset
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
class SystemParamService:
|
||||
CACHE_KEY = "sys:params"
|
||||
|
||||
def __init__(self, session: AsyncSession):
|
||||
self.session = session
|
||||
|
||||
async def get_value(self, code: str, *, from_cache: bool = True) -> str | None:
|
||||
if from_cache:
|
||||
try:
|
||||
cached = await java_hget(self.CACHE_KEY, code)
|
||||
if cached is not None:
|
||||
return str(cached)
|
||||
except Exception:
|
||||
logger.warning("Redis parameter cache read failed for %s", code, exc_info=True)
|
||||
result = await self.session.execute(
|
||||
text("SELECT param_value FROM sys_params WHERE param_code = :code LIMIT 1"),
|
||||
{"code": code},
|
||||
)
|
||||
value = result.scalar_one_or_none()
|
||||
if value is not None and from_cache:
|
||||
try:
|
||||
await java_hset(self.CACHE_KEY, code, str(value))
|
||||
except Exception:
|
||||
logger.warning("Redis parameter cache write failed for %s", code, exc_info=True)
|
||||
return None if value is None else str(value)
|
||||
|
||||
async def set_value(self, code: str, value: str) -> int:
|
||||
result = await self.session.execute(
|
||||
text(
|
||||
"UPDATE sys_params SET param_value = :value, update_date = CURRENT_TIMESTAMP WHERE param_code = :code"
|
||||
),
|
||||
{"code": code, "value": value},
|
||||
)
|
||||
try:
|
||||
await java_hset(self.CACHE_KEY, code, value)
|
||||
except Exception:
|
||||
logger.warning("Redis parameter cache write failed for %s", code, exc_info=True)
|
||||
return int(getattr(result, "rowcount", 0) or 0)
|
||||
@@ -0,0 +1,157 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from functools import lru_cache
|
||||
from typing import Any
|
||||
|
||||
from app.core.config import get_settings
|
||||
from app.core.i18n import _load_properties, message_for, resolve_language
|
||||
from app.core.ids import snowflake
|
||||
from app.core.redis import JavaRedisCodec, get_redis
|
||||
from app.core.security import AuthUser, shanghai_now_naive
|
||||
from app.repositories.timbre import TimbreRepository
|
||||
from app.schemas.timbre import TimbreBody
|
||||
|
||||
|
||||
def _details(row: dict[str, Any]) -> dict[str, Any]:
|
||||
return {
|
||||
"id": row.get("id"),
|
||||
"languages": row.get("languages"),
|
||||
"name": row.get("name"),
|
||||
"remark": row.get("remark"),
|
||||
"referenceAudio": row.get("reference_audio"),
|
||||
"referenceText": row.get("reference_text"),
|
||||
# TimbreDetailsVO.sort is primitive long, whose Java serializer always
|
||||
# emits a string and whose null conversion default is zero.
|
||||
"sort": str(row.get("sort") if row.get("sort") is not None else 0),
|
||||
"ttsModelId": row.get("tts_model_id"),
|
||||
"ttsVoice": row.get("tts_voice"),
|
||||
"voiceDemo": row.get("voice_demo"),
|
||||
}
|
||||
|
||||
|
||||
class TimbreService:
|
||||
def __init__(self, repository: TimbreRepository):
|
||||
self.repository = repository
|
||||
|
||||
@staticmethod
|
||||
def _validate(body: TimbreBody, language: str | None) -> None:
|
||||
from app.core.errors import AppError
|
||||
|
||||
for value, message in (
|
||||
(body.languages, "timbre.languages.require"),
|
||||
(body.name, "timbre.name.require"),
|
||||
(body.tts_model_id, "timbre.ttsModelId.require"),
|
||||
(body.tts_voice, "timbre.ttsVoice.require"),
|
||||
):
|
||||
if value is None or not value.strip():
|
||||
# TimbreController invokes ValidatorUtils directly. That
|
||||
# utility wraps validation text in RenException(String), whose
|
||||
# response code is 500 rather than the global @Valid code 10034.
|
||||
raise AppError(500, _validation_message(message, language))
|
||||
if body.sort is not None and body.sort < 0:
|
||||
raise AppError(500, _validation_message("sort.number", language))
|
||||
|
||||
async def page(
|
||||
self,
|
||||
tts_model_id: str | None,
|
||||
name: str | None,
|
||||
page: str | None,
|
||||
limit: str | None,
|
||||
language: str | None,
|
||||
) -> dict[str, Any]:
|
||||
if tts_model_id is None or not tts_model_id.strip():
|
||||
from app.core.errors import AppError
|
||||
|
||||
raise AppError(500, _validation_message("timbre.ttsModelId.require", language))
|
||||
current, size = max(int(page or "1"), 1), int(limit or "10")
|
||||
rows, total = await self.repository.page(
|
||||
tts_model_id=tts_model_id, name=name, offset=(current - 1) * size, limit=size
|
||||
)
|
||||
return {"total": total, "list": [_details(row) for row in rows]}
|
||||
|
||||
async def save(self, body: TimbreBody, user: AuthUser, language: str | None) -> None:
|
||||
self._validate(body, language)
|
||||
values = self._values(body, user, str(snowflake.next_id()))
|
||||
async with self.repository.session.begin():
|
||||
await self.repository.insert(values)
|
||||
|
||||
async def update(
|
||||
self, timbre_id: str, body: TimbreBody, user: AuthUser, language: str | None
|
||||
) -> None:
|
||||
self._validate(body, language)
|
||||
values = self._values(body, user, timbre_id)
|
||||
async with self.repository.session.begin():
|
||||
await self.repository.update(values)
|
||||
await get_redis().delete(f"timbre:details:{timbre_id}")
|
||||
|
||||
async def delete(self, ids: list[str]) -> None:
|
||||
async with self.repository.session.begin():
|
||||
await self.repository.delete(ids)
|
||||
|
||||
async def voices(self, model_id: str, voice_name: str | None, user: AuthUser, language: str | None) -> Any:
|
||||
normal, clones = await self.repository.voices(model_id, voice_name, user.id)
|
||||
values = [
|
||||
{
|
||||
"id": row.get("id"),
|
||||
"name": row.get("name"),
|
||||
"voiceDemo": row.get("voice_demo"),
|
||||
"languages": row.get("languages"),
|
||||
"isClone": False,
|
||||
}
|
||||
for row in normal
|
||||
]
|
||||
prefix = message_for(10158, language)
|
||||
redis = get_redis()
|
||||
for row in clones:
|
||||
name = prefix + str(row.get("name") or "")
|
||||
voice = {
|
||||
"id": row.get("id"),
|
||||
"name": name,
|
||||
"voiceDemo": row.get("voice_demo"),
|
||||
"languages": row.get("languages"),
|
||||
"isClone": True,
|
||||
}
|
||||
await redis.set(f"timbre:name:{row['id']}", JavaRedisCodec.encode(name))
|
||||
values.insert(0, voice)
|
||||
return values or None
|
||||
|
||||
@staticmethod
|
||||
def _values(body: TimbreBody, user: AuthUser, timbre_id: str) -> dict[str, Any]:
|
||||
assert body.languages is not None
|
||||
assert body.name is not None
|
||||
assert body.tts_model_id is not None
|
||||
assert body.tts_voice is not None
|
||||
return {
|
||||
"id": timbre_id,
|
||||
"languages": body.languages,
|
||||
"name": body.name,
|
||||
"remark": body.remark,
|
||||
"reference_audio": body.reference_audio,
|
||||
"reference_text": body.reference_text,
|
||||
"sort": body.sort if body.sort is not None else 0,
|
||||
"tts_model_id": body.tts_model_id,
|
||||
"tts_voice": body.tts_voice,
|
||||
"voice_demo": body.voice_demo,
|
||||
"creator": user.id,
|
||||
"updater": user.id,
|
||||
"now": shanghai_now_naive(),
|
||||
}
|
||||
|
||||
|
||||
_VALIDATION_FILES = {
|
||||
"zh-CN": "validation_zh_CN.properties",
|
||||
"zh-TW": "validation_zh_TW.properties",
|
||||
"en-US": "validation_en_US.properties",
|
||||
"de-DE": "validation_de_DE.properties",
|
||||
"vi-VN": "validation_vi_VN.properties",
|
||||
"pt-BR": "validation_pt_BR.properties",
|
||||
}
|
||||
|
||||
|
||||
@lru_cache(maxsize=64)
|
||||
def _validation_message(key: str, accept_language: str | None) -> str:
|
||||
language = resolve_language(accept_language)
|
||||
directory = get_settings().i18n_dir
|
||||
messages = _load_properties(directory / "validation.properties")
|
||||
messages.update(_load_properties(directory / _VALIDATION_FILES[language]))
|
||||
return messages.get(key, key)
|
||||
@@ -0,0 +1,334 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
import uuid
|
||||
from collections.abc import Mapping, Sequence
|
||||
from typing import Any
|
||||
|
||||
import httpx
|
||||
from redis.asyncio import Redis
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
|
||||
from app.core.config import get_settings
|
||||
from app.core.errors import AppError
|
||||
from app.core.i18n import message_for
|
||||
from app.core.security import AuthUser, shanghai_now_naive
|
||||
from app.integrations.voice_clone import VoiceCloneIntegration, VoiceCloneProviderError
|
||||
from app.repositories.voiceclone import VoiceCloneRepository
|
||||
from app.schemas.voiceclone import VoiceResourceCreateRequest
|
||||
from app.services.device import is_blank, redis_delete, redis_get, redis_set
|
||||
|
||||
VOICE_ORDER_COLUMNS = {
|
||||
"id": "id",
|
||||
"name": "name",
|
||||
"modelId": "model_id",
|
||||
"model_id": "model_id",
|
||||
"voiceId": "voice_id",
|
||||
"voice_id": "voice_id",
|
||||
"userId": "user_id",
|
||||
"user_id": "user_id",
|
||||
"trainStatus": "train_status",
|
||||
"train_status": "train_status",
|
||||
"createDate": "create_date",
|
||||
"create_date": "create_date",
|
||||
}
|
||||
|
||||
|
||||
class VoiceCloneService:
|
||||
def __init__(
|
||||
self,
|
||||
session: AsyncSession,
|
||||
*,
|
||||
redis_client: Redis | None = None,
|
||||
http_client: httpx.AsyncClient | None = None,
|
||||
provider: VoiceCloneIntegration | None = None,
|
||||
):
|
||||
self.session = session
|
||||
self.repository = VoiceCloneRepository(session)
|
||||
self.redis = redis_client
|
||||
self.provider = provider or VoiceCloneIntegration(
|
||||
timeout_seconds=get_settings().external_request_timeout_seconds,
|
||||
client=http_client,
|
||||
)
|
||||
|
||||
async def page(self, query: Mapping[str, Any], *, user_id: int | None = None) -> dict[str, Any]:
|
||||
page = int(str(query.get("page") or "1"))
|
||||
limit = int(str(query.get("limit") or "10"))
|
||||
name_value = query.get("name")
|
||||
name = None if name_value is None else str(name_value)
|
||||
effective_user = str(user_id) if user_id is not None else self._optional_string(query.get("userId"))
|
||||
requested = query.get("orderField")
|
||||
requested_fields = [requested] if isinstance(requested, str) else list(requested or [])
|
||||
order_fields = [VOICE_ORDER_COLUMNS[field] for field in requested_fields if field in VOICE_ORDER_COLUMNS]
|
||||
if not order_fields:
|
||||
order_fields = ["create_date"]
|
||||
ascending = str(query.get("order") or "").lower() == "asc" if requested_fields else True
|
||||
rows = await self.repository.page(
|
||||
page=page,
|
||||
limit=limit,
|
||||
name=name,
|
||||
user_id=effective_user,
|
||||
order_fields=order_fields,
|
||||
ascending=ascending,
|
||||
)
|
||||
return {
|
||||
"total": await self.repository.count(name=name, user_id=effective_user),
|
||||
"list": await self._response_list(rows),
|
||||
}
|
||||
|
||||
async def get_detail(self, voice_id: str) -> dict[str, Any] | None:
|
||||
row = await self.repository.get(voice_id)
|
||||
if row is None:
|
||||
return None
|
||||
return await self._response(row, include_has_voice=False)
|
||||
|
||||
async def get_by_user(self, user_id: int) -> list[dict[str, Any]]:
|
||||
del user_id
|
||||
# VoiceCloneServiceImpl.getByUserId orders ai_voice_clone by the
|
||||
# nonexistent ``created_at`` column (the schema uses ``create_date``).
|
||||
# The Java endpoint therefore consistently exposes its generic
|
||||
# code-500 envelope before result conversion.
|
||||
raise AppError(500)
|
||||
|
||||
async def create_resources(self, request: VoiceResourceCreateRequest, *, actor: AuthUser) -> None:
|
||||
model_id = request.model_id or ""
|
||||
config = await self._model_config(model_id)
|
||||
if config is None:
|
||||
raise AppError(10152)
|
||||
provider_type = config.get("type")
|
||||
if not isinstance(provider_type, str) or not provider_type.strip():
|
||||
raise AppError(10153)
|
||||
voice_ids = request.voice_ids or []
|
||||
for voice_id in voice_ids:
|
||||
if is_blank(voice_id):
|
||||
continue
|
||||
if provider_type == "huoshan_double_stream" and "S_" not in voice_id:
|
||||
raise AppError(10160)
|
||||
if await self.repository.voice_id_count(model_id=model_id, voice_id=voice_id):
|
||||
raise AppError(10159)
|
||||
|
||||
now = shanghai_now_naive()
|
||||
prefix = now.strftime("%m%d%H%M")
|
||||
values: list[dict[str, Any]] = []
|
||||
for index, voice_id in enumerate(voice_ids, start=1):
|
||||
values.append(
|
||||
{
|
||||
"id": uuid.uuid4().hex,
|
||||
"name": f"{prefix}_{index}",
|
||||
"model_id": model_id,
|
||||
"voice_id": voice_id,
|
||||
"languages": request.languages,
|
||||
"user_id": request.user_id,
|
||||
"voice": None,
|
||||
"train_status": 0,
|
||||
"train_error": None,
|
||||
"creator": actor.id,
|
||||
"create_date": now,
|
||||
}
|
||||
)
|
||||
try:
|
||||
await self.repository.insert_many(values)
|
||||
await self.session.commit()
|
||||
except Exception:
|
||||
await self.session.rollback()
|
||||
raise
|
||||
|
||||
async def delete(self, ids: Sequence[str]) -> None:
|
||||
await self.repository.delete_many(ids)
|
||||
await self.session.commit()
|
||||
|
||||
async def check_permission(self, voice_id: str | None, user: AuthUser) -> dict[str, Any]:
|
||||
row = await self.repository.get(voice_id)
|
||||
if row is None:
|
||||
raise AppError(10144)
|
||||
if int(row.get("user_id") or -1) != user.id:
|
||||
raise AppError(10150)
|
||||
return row
|
||||
|
||||
async def upload_voice(self, voice_id: str, content: bytes) -> None:
|
||||
if await self.repository.get(voice_id) is None:
|
||||
raise AppError(10144)
|
||||
await self.repository.update_voice(voice_id, content)
|
||||
await self.session.commit()
|
||||
|
||||
async def rename(self, voice_id: str, name: str) -> None:
|
||||
if await self.repository.get(voice_id) is None:
|
||||
raise AppError(10144)
|
||||
await self.repository.update_name(voice_id, name)
|
||||
await self.session.commit()
|
||||
await redis_delete(f"timbre:name:{voice_id}", client=self.redis)
|
||||
|
||||
async def create_audio_id(self, voice_id: str) -> str:
|
||||
row = await self.repository.get(voice_id)
|
||||
if row is None or row.get("voice") is None:
|
||||
raise AppError(10182)
|
||||
value = str(uuid.uuid4())
|
||||
await redis_set(f"voiceClone:audio:id:{value}", voice_id, client=self.redis)
|
||||
return value
|
||||
|
||||
async def consume_audio(self, download_id: str) -> bytes | None:
|
||||
key = f"voiceClone:audio:id:{download_id}"
|
||||
voice_id = await redis_get(key, self.redis)
|
||||
await redis_delete(key, client=self.redis)
|
||||
if is_blank(None if voice_id is None else str(voice_id)):
|
||||
return None
|
||||
row = await self.repository.get(str(voice_id))
|
||||
data = None if row is None else row.get("voice")
|
||||
if data is None:
|
||||
return None
|
||||
result = bytes(data)
|
||||
return result or None
|
||||
|
||||
async def clone_audio(
|
||||
self,
|
||||
voice_id: str,
|
||||
*,
|
||||
accept_language: str | None,
|
||||
) -> None:
|
||||
row = await self.repository.get(voice_id)
|
||||
if row is None:
|
||||
raise AppError(10144)
|
||||
raw_voice = row.get("voice")
|
||||
if raw_voice is None or len(raw_voice) == 0:
|
||||
raise AppError(10151)
|
||||
try:
|
||||
config = await self._model_config(str(row.get("model_id") or ""))
|
||||
if config is None:
|
||||
raise AppError(10152)
|
||||
provider_type = config.get("type")
|
||||
if not isinstance(provider_type, str) or not provider_type.strip():
|
||||
raise AppError(10153)
|
||||
if provider_type != "huoshan_double_stream":
|
||||
return
|
||||
appid = config.get("appid")
|
||||
access_token = config.get("access_token")
|
||||
if (
|
||||
not isinstance(appid, str)
|
||||
or is_blank(appid)
|
||||
or not isinstance(access_token, str)
|
||||
or is_blank(access_token)
|
||||
):
|
||||
raise AppError(10155)
|
||||
speaker_id = await self.provider.train_huoshan(
|
||||
appid=appid,
|
||||
access_token=access_token,
|
||||
voice=bytes(raw_voice),
|
||||
speaker_id=str(row.get("voice_id") or ""),
|
||||
)
|
||||
await self.repository.update_training(
|
||||
voice_id,
|
||||
train_status=2,
|
||||
train_error="",
|
||||
speaker_id=speaker_id,
|
||||
)
|
||||
await self.session.commit()
|
||||
except AppError as exc:
|
||||
await self._record_training_failure(voice_id, exc.message or message_for(exc.code, accept_language))
|
||||
raise
|
||||
except VoiceCloneProviderError as exc:
|
||||
if exc.code in {500, 10156}:
|
||||
await self._record_training_failure(voice_id, exc.message)
|
||||
raise AppError(exc.code, exc.message) from exc
|
||||
translated = message_for(10154, accept_language, exc.message)
|
||||
await self._record_training_failure(voice_id, translated)
|
||||
raise AppError(10154, translated) from exc
|
||||
except Exception as exc:
|
||||
translated = message_for(10154, accept_language, str(exc))
|
||||
await self._record_training_failure(voice_id, translated)
|
||||
raise AppError(10154, translated) from exc
|
||||
|
||||
async def tts_platforms(self) -> list[dict[str, Any]]:
|
||||
return await self.repository.get_tts_platforms()
|
||||
|
||||
async def _record_training_failure(self, voice_id: str, message: str) -> None:
|
||||
await self.session.rollback()
|
||||
await self.repository.update_training(voice_id, train_status=3, train_error=message)
|
||||
await self.session.commit()
|
||||
|
||||
async def _model_config(self, model_id: str) -> dict[str, Any] | None:
|
||||
if is_blank(model_id):
|
||||
return None
|
||||
cached = await redis_get(f"model:data:{model_id}", self.redis)
|
||||
cached_mapping = self._mapping(cached)
|
||||
if cached_mapping is not None:
|
||||
config_value = cached_mapping.get("configJson", cached_mapping.get("config_json"))
|
||||
parsed = self._json_mapping(config_value)
|
||||
if parsed is not None:
|
||||
return parsed
|
||||
row = await self.repository.get_model_config(model_id)
|
||||
return None if row is None else self._json_mapping(row.get("config_json"))
|
||||
|
||||
async def _model_name(self, model_id: str | None) -> str | None:
|
||||
if is_blank(model_id):
|
||||
return None
|
||||
cache_key = f"model:name:{model_id}"
|
||||
cached = await redis_get(cache_key, self.redis)
|
||||
if isinstance(cached, str) and cached.strip():
|
||||
return cached
|
||||
value = await self.repository.get_model_name(model_id or "")
|
||||
if value is not None and value.strip():
|
||||
await redis_set(cache_key, value, client=self.redis)
|
||||
return value
|
||||
|
||||
async def _response_list(self, rows: Sequence[Mapping[str, Any]]) -> list[dict[str, Any]]:
|
||||
user_ids = [int(row["user_id"]) for row in rows if row.get("user_id") is not None]
|
||||
usernames = await self.repository.get_usernames(user_ids)
|
||||
result: list[dict[str, Any]] = []
|
||||
for row in rows:
|
||||
result.append(await self._response(row, usernames=usernames, include_has_voice=True))
|
||||
return result
|
||||
|
||||
async def _response(
|
||||
self,
|
||||
row: Mapping[str, Any],
|
||||
*,
|
||||
usernames: Mapping[int, str] | None = None,
|
||||
include_has_voice: bool,
|
||||
) -> dict[str, Any]:
|
||||
user_id = None if row.get("user_id") is None else int(row["user_id"])
|
||||
if user_id is None:
|
||||
username = None
|
||||
elif usernames is None:
|
||||
username = await self.repository.get_username(user_id)
|
||||
else:
|
||||
username = usernames.get(user_id)
|
||||
return {
|
||||
"id": row.get("id"),
|
||||
"name": row.get("name"),
|
||||
"model_id": row.get("model_id"),
|
||||
"model_name": await self._model_name(self._optional_string(row.get("model_id"))),
|
||||
"voice_id": row.get("voice_id"),
|
||||
"languages": row.get("languages"),
|
||||
"user_id": user_id,
|
||||
"user_name": username,
|
||||
"train_status": row.get("train_status"),
|
||||
"train_error": row.get("train_error"),
|
||||
"create_date": row.get("create_date"),
|
||||
"has_voice": row.get("voice") is not None if include_has_voice else None,
|
||||
}
|
||||
|
||||
@staticmethod
|
||||
def _mapping(value: Any) -> dict[str, Any] | None:
|
||||
if isinstance(value, dict):
|
||||
return {str(key): item for key, item in value.items() if key != "@class"}
|
||||
if isinstance(value, list) and len(value) == 2 and isinstance(value[1], dict):
|
||||
return {str(key): item for key, item in value[1].items()}
|
||||
return None
|
||||
|
||||
@staticmethod
|
||||
def _json_mapping(value: Any) -> dict[str, Any] | None:
|
||||
if isinstance(value, dict):
|
||||
return {str(key): item for key, item in value.items()}
|
||||
if isinstance(value, bytes):
|
||||
value = value.decode("utf-8")
|
||||
if isinstance(value, str):
|
||||
try:
|
||||
parsed = json.loads(value)
|
||||
except json.JSONDecodeError:
|
||||
return None
|
||||
return {str(key): item for key, item in parsed.items()} if isinstance(parsed, dict) else None
|
||||
return None
|
||||
|
||||
@staticmethod
|
||||
def _optional_string(value: Any) -> str | None:
|
||||
return None if value is None else str(value)
|
||||
File diff suppressed because it is too large
Load Diff
File diff suppressed because it is too large
Load Diff
File diff suppressed because it is too large
Load Diff
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,164 @@
|
||||
{
|
||||
"schema_version": 1,
|
||||
"generated_at": "2026-07-20T07:12:33.158826+00:00",
|
||||
"parameters": {
|
||||
"requests_per_service_per_scenario": 60,
|
||||
"concurrency": 6,
|
||||
"sequential_warmup_requests": 10,
|
||||
"scenario_order": [
|
||||
"representative-read",
|
||||
"representative-crud-update",
|
||||
"runtime-configuration",
|
||||
"ota-check-and-signing"
|
||||
],
|
||||
"service_order": [
|
||||
"java",
|
||||
"fastapi"
|
||||
]
|
||||
},
|
||||
"results": [
|
||||
{
|
||||
"service": "java",
|
||||
"scenario": "representative-read",
|
||||
"requests": 60,
|
||||
"concurrency": 6,
|
||||
"warmup_requests": 10,
|
||||
"errors": 0,
|
||||
"elapsed_seconds": 0.099812,
|
||||
"throughput_requests_per_second": 601.128,
|
||||
"latency_ms_min": 3.259,
|
||||
"latency_ms_p50": 7.662,
|
||||
"latency_ms_p95": 16.239,
|
||||
"latency_ms_max": 20.352
|
||||
},
|
||||
{
|
||||
"service": "fastapi",
|
||||
"scenario": "representative-read",
|
||||
"requests": 60,
|
||||
"concurrency": 6,
|
||||
"warmup_requests": 10,
|
||||
"errors": 0,
|
||||
"elapsed_seconds": 0.079836,
|
||||
"throughput_requests_per_second": 751.538,
|
||||
"latency_ms_min": 4.117,
|
||||
"latency_ms_p50": 6.749,
|
||||
"latency_ms_p95": 12.552,
|
||||
"latency_ms_max": 16.057
|
||||
},
|
||||
{
|
||||
"service": "java",
|
||||
"scenario": "representative-crud-update",
|
||||
"requests": 60,
|
||||
"concurrency": 6,
|
||||
"warmup_requests": 10,
|
||||
"errors": 0,
|
||||
"elapsed_seconds": 0.110495,
|
||||
"throughput_requests_per_second": 543.013,
|
||||
"latency_ms_min": 5.684,
|
||||
"latency_ms_p50": 9.331,
|
||||
"latency_ms_p95": 15.302,
|
||||
"latency_ms_max": 19.008
|
||||
},
|
||||
{
|
||||
"service": "fastapi",
|
||||
"scenario": "representative-crud-update",
|
||||
"requests": 60,
|
||||
"concurrency": 6,
|
||||
"warmup_requests": 10,
|
||||
"errors": 0,
|
||||
"elapsed_seconds": 0.12582,
|
||||
"throughput_requests_per_second": 476.87,
|
||||
"latency_ms_min": 6.996,
|
||||
"latency_ms_p50": 11.252,
|
||||
"latency_ms_p95": 20.599,
|
||||
"latency_ms_max": 24.284
|
||||
},
|
||||
{
|
||||
"service": "java",
|
||||
"scenario": "runtime-configuration",
|
||||
"requests": 60,
|
||||
"concurrency": 6,
|
||||
"warmup_requests": 10,
|
||||
"errors": 0,
|
||||
"elapsed_seconds": 0.105474,
|
||||
"throughput_requests_per_second": 568.863,
|
||||
"latency_ms_min": 3.901,
|
||||
"latency_ms_p50": 8.746,
|
||||
"latency_ms_p95": 16.127,
|
||||
"latency_ms_max": 29.528
|
||||
},
|
||||
{
|
||||
"service": "fastapi",
|
||||
"scenario": "runtime-configuration",
|
||||
"requests": 60,
|
||||
"concurrency": 6,
|
||||
"warmup_requests": 10,
|
||||
"errors": 0,
|
||||
"elapsed_seconds": 0.078418,
|
||||
"throughput_requests_per_second": 765.126,
|
||||
"latency_ms_min": 3.868,
|
||||
"latency_ms_p50": 6.488,
|
||||
"latency_ms_p95": 13.722,
|
||||
"latency_ms_max": 18.923
|
||||
},
|
||||
{
|
||||
"service": "java",
|
||||
"scenario": "ota-check-and-signing",
|
||||
"requests": 60,
|
||||
"concurrency": 6,
|
||||
"warmup_requests": 10,
|
||||
"errors": 0,
|
||||
"elapsed_seconds": 0.131968,
|
||||
"throughput_requests_per_second": 454.655,
|
||||
"latency_ms_min": 7.543,
|
||||
"latency_ms_p50": 11.301,
|
||||
"latency_ms_p95": 19.854,
|
||||
"latency_ms_max": 21.729
|
||||
},
|
||||
{
|
||||
"service": "fastapi",
|
||||
"scenario": "ota-check-and-signing",
|
||||
"requests": 60,
|
||||
"concurrency": 6,
|
||||
"warmup_requests": 10,
|
||||
"errors": 0,
|
||||
"elapsed_seconds": 0.169106,
|
||||
"throughput_requests_per_second": 354.807,
|
||||
"latency_ms_min": 7.066,
|
||||
"latency_ms_p50": 16.208,
|
||||
"latency_ms_p95": 20.344,
|
||||
"latency_ms_max": 28.953
|
||||
}
|
||||
],
|
||||
"comparisons": [
|
||||
{
|
||||
"scenario": "representative-read",
|
||||
"p50_ratio_fastapi_over_java": 0.881,
|
||||
"p95_ratio_fastapi_over_java": 0.773,
|
||||
"throughput_ratio_fastapi_over_java": 1.25
|
||||
},
|
||||
{
|
||||
"scenario": "representative-crud-update",
|
||||
"p50_ratio_fastapi_over_java": 1.206,
|
||||
"p95_ratio_fastapi_over_java": 1.346,
|
||||
"throughput_ratio_fastapi_over_java": 0.878
|
||||
},
|
||||
{
|
||||
"scenario": "runtime-configuration",
|
||||
"p50_ratio_fastapi_over_java": 0.742,
|
||||
"p95_ratio_fastapi_over_java": 0.851,
|
||||
"throughput_ratio_fastapi_over_java": 1.345
|
||||
},
|
||||
{
|
||||
"scenario": "ota-check-and-signing",
|
||||
"p50_ratio_fastapi_over_java": 1.434,
|
||||
"p95_ratio_fastapi_over_java": 1.025,
|
||||
"throughput_ratio_fastapi_over_java": 0.78
|
||||
}
|
||||
],
|
||||
"summary": {
|
||||
"measurements": 8,
|
||||
"requests_measured": 480,
|
||||
"errors": 0
|
||||
}
|
||||
}
|
||||
File diff suppressed because it is too large
Load Diff
+16
@@ -0,0 +1,16 @@
|
||||
#!/bin/sh
|
||||
set -eu
|
||||
|
||||
UPSTREAM=${MANAGER_API_UPSTREAM:-manager-api-fastapi:8002}
|
||||
case "${UPSTREAM}" in
|
||||
''|*[!A-Za-z0-9._:-]*)
|
||||
echo "MANAGER_API_UPSTREAM must be a hostname-or-IP and port" >&2
|
||||
exit 2
|
||||
;;
|
||||
esac
|
||||
|
||||
export MANAGER_API_UPSTREAM=${UPSTREAM}
|
||||
envsubst '${MANAGER_API_UPSTREAM}' \
|
||||
< /etc/nginx/nginx.conf.template \
|
||||
> /tmp/manager-api-nginx.conf
|
||||
exec nginx -c /tmp/manager-api-nginx.conf -g 'daemon off;'
|
||||
@@ -0,0 +1,45 @@
|
||||
worker_processes auto;
|
||||
pid /tmp/nginx.pid;
|
||||
|
||||
events {
|
||||
worker_connections 1024;
|
||||
}
|
||||
|
||||
http {
|
||||
include /etc/nginx/mime.types;
|
||||
default_type application/octet-stream;
|
||||
access_log /dev/stdout;
|
||||
error_log /dev/stderr warn;
|
||||
sendfile on;
|
||||
keepalive_timeout 65;
|
||||
client_max_body_size 100m;
|
||||
|
||||
upstream manager_api_fastapi {
|
||||
server ${MANAGER_API_UPSTREAM};
|
||||
keepalive 32;
|
||||
}
|
||||
|
||||
server {
|
||||
listen 8002;
|
||||
server_name _;
|
||||
|
||||
location = /xiaozhi {
|
||||
return 308 /xiaozhi/;
|
||||
}
|
||||
|
||||
location /xiaozhi/ {
|
||||
proxy_pass http://manager_api_fastapi;
|
||||
proxy_http_version 1.1;
|
||||
proxy_set_header Host $host;
|
||||
proxy_set_header X-Real-IP $remote_addr;
|
||||
proxy_set_header X-Forwarded-For $proxy_add_x_forwarded_for;
|
||||
proxy_set_header X-Forwarded-Proto $scheme;
|
||||
proxy_set_header Connection "";
|
||||
proxy_connect_timeout 10s;
|
||||
proxy_send_timeout 130s;
|
||||
proxy_read_timeout 130s;
|
||||
proxy_request_buffering off;
|
||||
proxy_buffering off;
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,93 @@
|
||||
services:
|
||||
manager-api-migrate:
|
||||
image: xiaozhi/manager-api-migrate:fastapi-0.1.0
|
||||
build:
|
||||
context: ../..
|
||||
dockerfile: main/manager-api-fastapi/Dockerfile.migrations
|
||||
environment:
|
||||
LIQUIBASE_URL: ${LIQUIBASE_URL:?Set a JDBC MySQL URL}
|
||||
LIQUIBASE_USERNAME: ${MYSQL_USER:?Set MYSQL_USER}
|
||||
LIQUIBASE_PASSWORD: ${MYSQL_PASSWORD:?Set MYSQL_PASSWORD}
|
||||
MIGRATION_POM: /migration/pom.xml
|
||||
JAVA_RESOURCES_DIR: /migration/java-resources
|
||||
MAVEN_BIN: mvn
|
||||
restart: "no"
|
||||
|
||||
manager-api-fastapi:
|
||||
image: xiaozhi/manager-api-fastapi:0.1.0
|
||||
build:
|
||||
context: ../..
|
||||
dockerfile: main/manager-api-fastapi/Dockerfile
|
||||
init: true
|
||||
depends_on:
|
||||
manager-api-migrate:
|
||||
condition: service_completed_successfully
|
||||
environment:
|
||||
APP_ENVIRONMENT: production
|
||||
APP_DATABASE_URL: ${FASTAPI_DATABASE_URL:?Set an asyncmy MySQL URL}
|
||||
APP_REDIS_URL: ${REDIS_URL:?Set REDIS_URL}
|
||||
APP_WORKERS: ${APP_WORKERS:-2}
|
||||
APP_GRACEFUL_SHUTDOWN_SECONDS: ${APP_GRACEFUL_SHUTDOWN_SECONDS:-30}
|
||||
APP_FORWARDED_ALLOW_IPS: ${APP_FORWARDED_ALLOW_IPS:-*}
|
||||
APP_ALLOW_START_WITHOUT_DEPENDENCIES: "false"
|
||||
expose:
|
||||
- "8002"
|
||||
volumes:
|
||||
# During cutover this source can be the retained Java service's host
|
||||
# uploadfile directory so both implementations see identical files.
|
||||
- ${MANAGER_API_UPLOAD_SOURCE:-manager-api-uploads}:/data/uploads
|
||||
read_only: true
|
||||
tmpfs:
|
||||
- /tmp:size=64m,mode=1777
|
||||
restart: unless-stopped
|
||||
stop_grace_period: 40s
|
||||
healthcheck:
|
||||
test: ["CMD", "python", "-c", "import urllib.request; urllib.request.urlopen('http://127.0.0.1:8002/xiaozhi/health/ready', timeout=2).read()"]
|
||||
interval: 15s
|
||||
timeout: 3s
|
||||
retries: 4
|
||||
start_period: 20s
|
||||
|
||||
manager-api-jobs:
|
||||
image: xiaozhi/manager-api-fastapi:0.1.0
|
||||
init: true
|
||||
depends_on:
|
||||
manager-api-migrate:
|
||||
condition: service_completed_successfully
|
||||
command: ["python", "-m", "app.jobs.worker"]
|
||||
environment:
|
||||
APP_ENVIRONMENT: production
|
||||
APP_DATABASE_URL: ${FASTAPI_DATABASE_URL:?Set an asyncmy MySQL URL}
|
||||
APP_REDIS_URL: ${REDIS_URL:?Set REDIS_URL}
|
||||
APP_GRACEFUL_SHUTDOWN_SECONDS: ${APP_GRACEFUL_SHUTDOWN_SECONDS:-30}
|
||||
APP_ALLOW_START_WITHOUT_DEPENDENCIES: "false"
|
||||
volumes:
|
||||
- ${MANAGER_API_UPLOAD_SOURCE:-manager-api-uploads}:/data/uploads
|
||||
read_only: true
|
||||
tmpfs:
|
||||
- /tmp:size=64m,mode=1777
|
||||
restart: unless-stopped
|
||||
stop_grace_period: 40s
|
||||
|
||||
manager-api-nginx:
|
||||
image: xiaozhi/manager-api-nginx:fastapi-0.1.0
|
||||
build:
|
||||
context: ../..
|
||||
dockerfile: main/manager-api-fastapi/Dockerfile.nginx
|
||||
depends_on:
|
||||
manager-api-fastapi:
|
||||
condition: service_healthy
|
||||
ports:
|
||||
- "8002:8002"
|
||||
environment:
|
||||
MANAGER_API_UPSTREAM: ${MANAGER_API_UPSTREAM:-manager-api-fastapi:8002}
|
||||
read_only: true
|
||||
tmpfs:
|
||||
- /var/cache/nginx:size=32m
|
||||
- /var/run:size=1m
|
||||
- /tmp:size=16m
|
||||
restart: unless-stopped
|
||||
stop_grace_period: 15s
|
||||
|
||||
volumes:
|
||||
manager-api-uploads:
|
||||
@@ -0,0 +1,84 @@
|
||||
<?xml version="1.0" encoding="UTF-8"?>
|
||||
<project xmlns="http://maven.apache.org/POM/4.0.0"
|
||||
xmlns:xsi="http://www.w3.org/2001/XMLSchema-instance"
|
||||
xsi:schemaLocation="http://maven.apache.org/POM/4.0.0 https://maven.apache.org/xsd/maven-4.0.0.xsd">
|
||||
<modelVersion>4.0.0</modelVersion>
|
||||
<groupId>xiaozhi</groupId>
|
||||
<artifactId>manager-api-liquibase-runner</artifactId>
|
||||
<version>1.0.0</version>
|
||||
<properties>
|
||||
<maven.compiler.release>21</maven.compiler.release>
|
||||
<project.build.sourceEncoding>UTF-8</project.build.sourceEncoding>
|
||||
<java.resources.dir>${project.basedir}/../manager-api/src/main/resources</java.resources.dir>
|
||||
<liquibase.version>4.20.0</liquibase.version>
|
||||
<mysql.version>9.1.0</mysql.version>
|
||||
<spring.version>6.2.3</spring.version>
|
||||
</properties>
|
||||
<dependencies>
|
||||
<dependency>
|
||||
<groupId>org.liquibase</groupId>
|
||||
<artifactId>liquibase-core</artifactId>
|
||||
<version>${liquibase.version}</version>
|
||||
</dependency>
|
||||
<dependency>
|
||||
<groupId>org.springframework</groupId>
|
||||
<artifactId>spring-jdbc</artifactId>
|
||||
<version>${spring.version}</version>
|
||||
</dependency>
|
||||
<dependency>
|
||||
<groupId>org.springframework</groupId>
|
||||
<artifactId>spring-context</artifactId>
|
||||
<version>${spring.version}</version>
|
||||
</dependency>
|
||||
<dependency>
|
||||
<groupId>com.mysql</groupId>
|
||||
<artifactId>mysql-connector-j</artifactId>
|
||||
<version>${mysql.version}</version>
|
||||
</dependency>
|
||||
<dependency>
|
||||
<groupId>org.slf4j</groupId>
|
||||
<artifactId>slf4j-simple</artifactId>
|
||||
<version>2.0.16</version>
|
||||
</dependency>
|
||||
</dependencies>
|
||||
<build>
|
||||
<sourceDirectory>${project.basedir}/migration-src</sourceDirectory>
|
||||
<resources>
|
||||
<resource>
|
||||
<directory>${java.resources.dir}</directory>
|
||||
<filtering>false</filtering>
|
||||
</resource>
|
||||
</resources>
|
||||
<plugins>
|
||||
<plugin>
|
||||
<groupId>org.apache.maven.plugins</groupId>
|
||||
<artifactId>maven-compiler-plugin</artifactId>
|
||||
<version>3.13.0</version>
|
||||
</plugin>
|
||||
<plugin>
|
||||
<groupId>org.apache.maven.plugins</groupId>
|
||||
<artifactId>maven-shade-plugin</artifactId>
|
||||
<version>3.6.0</version>
|
||||
<executions>
|
||||
<execution>
|
||||
<phase>package</phase>
|
||||
<goals>
|
||||
<goal>shade</goal>
|
||||
</goals>
|
||||
<configuration>
|
||||
<createDependencyReducedPom>false</createDependencyReducedPom>
|
||||
<shadedArtifactAttached>true</shadedArtifactAttached>
|
||||
<shadedClassifierName>all</shadedClassifierName>
|
||||
<transformers>
|
||||
<transformer implementation="org.apache.maven.plugins.shade.resource.ManifestResourceTransformer">
|
||||
<mainClass>xiaozhi.migration.LiquibaseMigrationRunner</mainClass>
|
||||
</transformer>
|
||||
<transformer implementation="org.apache.maven.plugins.shade.resource.ServicesResourceTransformer"/>
|
||||
</transformers>
|
||||
</configuration>
|
||||
</execution>
|
||||
</executions>
|
||||
</plugin>
|
||||
</plugins>
|
||||
</build>
|
||||
</project>
|
||||
+57
@@ -0,0 +1,57 @@
|
||||
package xiaozhi.migration;
|
||||
|
||||
import java.sql.Connection;
|
||||
import java.sql.ResultSet;
|
||||
import java.sql.Statement;
|
||||
|
||||
import javax.sql.DataSource;
|
||||
|
||||
import liquibase.integration.spring.SpringLiquibase;
|
||||
import org.springframework.core.io.DefaultResourceLoader;
|
||||
import org.springframework.jdbc.datasource.DriverManagerDataSource;
|
||||
|
||||
/** Runs only the original Spring Liquibase changelog; it never starts manager-api or Redis. */
|
||||
public final class LiquibaseMigrationRunner {
|
||||
private LiquibaseMigrationRunner() {
|
||||
}
|
||||
|
||||
public static void main(String[] args) throws Exception {
|
||||
String url = requiredEnvironment("MIGRATION_JDBC_URL");
|
||||
String username = requiredEnvironment("MIGRATION_USERNAME");
|
||||
String password = requiredEnvironment("MIGRATION_PASSWORD");
|
||||
|
||||
DriverManagerDataSource dataSource = new DriverManagerDataSource(url, username, password);
|
||||
dataSource.setDriverClassName("com.mysql.cj.jdbc.Driver");
|
||||
runLiquibase(dataSource);
|
||||
System.out.println("Liquibase migration complete; applied changeSets=" + appliedChangeSetCount(dataSource));
|
||||
}
|
||||
|
||||
private static void runLiquibase(DataSource dataSource) throws Exception {
|
||||
SpringLiquibase liquibase = new SpringLiquibase();
|
||||
liquibase.setDataSource(dataSource);
|
||||
liquibase.setChangeLog("classpath:db/changelog/db.changelog-master.yaml");
|
||||
liquibase.setResourceLoader(new DefaultResourceLoader(Thread.currentThread().getContextClassLoader()));
|
||||
liquibase.setDropFirst(false);
|
||||
liquibase.setShouldRun(true);
|
||||
liquibase.afterPropertiesSet();
|
||||
}
|
||||
|
||||
private static int appliedChangeSetCount(DataSource dataSource) throws Exception {
|
||||
try (Connection connection = dataSource.getConnection();
|
||||
Statement statement = connection.createStatement();
|
||||
ResultSet result = statement.executeQuery("SELECT COUNT(*) FROM DATABASECHANGELOG")) {
|
||||
if (!result.next()) {
|
||||
throw new IllegalStateException("DATABASECHANGELOG count query returned no row");
|
||||
}
|
||||
return result.getInt(1);
|
||||
}
|
||||
}
|
||||
|
||||
private static String requiredEnvironment(String name) {
|
||||
String value = System.getenv(name);
|
||||
if (value == null || value.isBlank()) {
|
||||
throw new IllegalArgumentException("Required environment variable is missing: " + name);
|
||||
}
|
||||
return value;
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,75 @@
|
||||
[project]
|
||||
name = "xiaozhi-manager-api-fastapi"
|
||||
version = "0.1.0"
|
||||
description = "FastAPI-compatible replacement for xiaozhi manager-api"
|
||||
readme = "README.md"
|
||||
requires-python = ">=3.10,<3.13"
|
||||
dependencies = [
|
||||
"aiosqlite==0.21.0",
|
||||
"asyncmy==0.2.10",
|
||||
"bcrypt==4.3.0",
|
||||
"cryptography==45.0.5",
|
||||
"fastapi==0.116.1",
|
||||
"gmssl==3.2.2",
|
||||
"httpx==0.28.1",
|
||||
"pillow==11.3.0",
|
||||
# FastAPI 0.116 reconstructs request fields through TypeAdapter. Pydantic
|
||||
# 2.13 warns that aliases on that compatibility path are ineffective; pin
|
||||
# the contemporary 2.11 line so request aliases remain warning-free.
|
||||
"pydantic==2.11.7",
|
||||
"pydantic-settings==2.10.1",
|
||||
"python-multipart==0.0.20",
|
||||
"pyyaml==6.0.2",
|
||||
"redis[hiredis]==6.2.0",
|
||||
"sqlalchemy[asyncio]==2.0.41",
|
||||
"uvicorn[standard]==0.35.0",
|
||||
"websockets==15.0.1",
|
||||
]
|
||||
|
||||
[dependency-groups]
|
||||
dev = [
|
||||
"mypy==1.17.0",
|
||||
"pytest==8.4.1",
|
||||
"pytest-asyncio==1.1.0",
|
||||
"pytest-cov==6.2.1",
|
||||
"respx==0.22.0",
|
||||
"ruff==0.12.4",
|
||||
]
|
||||
|
||||
[tool.uv]
|
||||
default-groups = ["dev"]
|
||||
|
||||
[tool.pytest.ini_options]
|
||||
addopts = "-ra --strict-config --strict-markers"
|
||||
asyncio_mode = "auto"
|
||||
testpaths = ["tests"]
|
||||
markers = [
|
||||
"integration: requires isolated MySQL and Redis",
|
||||
"contract: compares the Java and FastAPI services",
|
||||
"performance: runs the representative performance comparison",
|
||||
]
|
||||
|
||||
[tool.ruff]
|
||||
target-version = "py310"
|
||||
line-length = 120
|
||||
exclude = ["tests/fixtures/generated"]
|
||||
|
||||
[tool.ruff.lint]
|
||||
select = ["E", "F", "I", "UP", "B", "ASYNC", "S"]
|
||||
ignore = ["S101"]
|
||||
|
||||
[tool.ruff.lint.per-file-ignores]
|
||||
"app/routers/*.py" = ["B008"]
|
||||
|
||||
[tool.mypy]
|
||||
python_version = "3.10"
|
||||
strict = true
|
||||
plugins = ["pydantic.mypy"]
|
||||
exclude = ["tests/fixtures/generated"]
|
||||
|
||||
[build-system]
|
||||
requires = ["hatchling==1.27.0"]
|
||||
build-backend = "hatchling.build"
|
||||
|
||||
[tool.hatch.build.targets.wheel]
|
||||
packages = ["app"]
|
||||
@@ -0,0 +1,14 @@
|
||||
#!/bin/sh
|
||||
set -eu
|
||||
|
||||
if [ "$#" -gt 0 ]; then
|
||||
exec "$@"
|
||||
fi
|
||||
|
||||
exec uvicorn app.main:app \
|
||||
--host "${APP_HOST:-0.0.0.0}" \
|
||||
--port "${APP_PORT:-8002}" \
|
||||
--workers "${APP_WORKERS:-2}" \
|
||||
--timeout-graceful-shutdown "${APP_GRACEFUL_SHUTDOWN_SECONDS:-30}" \
|
||||
--proxy-headers \
|
||||
--forwarded-allow-ips "${APP_FORWARDED_ALLOW_IPS:-127.0.0.1}"
|
||||
@@ -0,0 +1,267 @@
|
||||
#!/usr/bin/env python3
|
||||
"""Extract manager-api HTTP call sites from all three in-repository consumers."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
import re
|
||||
from collections import Counter
|
||||
from pathlib import Path
|
||||
from typing import NamedTuple
|
||||
|
||||
TARGET_ROOT = Path(__file__).resolve().parents[1]
|
||||
REPO_ROOT = TARGET_ROOT.parents[1]
|
||||
MAIN_ROOT = REPO_ROOT / "main"
|
||||
|
||||
WEB_ROOT = MAIN_ROOT / "manager-web" / "src"
|
||||
MOBILE_ROOT = MAIN_ROOT / "manager-mobile" / "src"
|
||||
SERVER_ROOT = MAIN_ROOT / "xiaozhi-server"
|
||||
|
||||
WEB_CHAIN = re.compile(
|
||||
r"\.url\(\s*(?P<quote>[`'\"])(?P<url>.*?)(?P=quote)\s*\)"
|
||||
r"(?:\s*//[^\n]*)?\s*\.method\(\s*['\"](?P<method>[A-Za-z]+)['\"]\s*\)",
|
||||
re.DOTALL,
|
||||
)
|
||||
WEB_CONFIG = re.compile(
|
||||
r"\burl\s*:\s*(?P<quote>[`'\"])(?P<url>.*?)(?P=quote)\s*,"
|
||||
r"\s*method\s*:\s*['\"](?P<method>[A-Za-z]+)['\"]",
|
||||
re.DOTALL,
|
||||
)
|
||||
WEB_TEMPLATE_URL = re.compile(
|
||||
r"(?P<quote>`)(?P<url>\$\{(?:(?:Api|api)\.)?getServiceUrl\(\)\}/.*?)"
|
||||
r"(?P=quote)"
|
||||
)
|
||||
WEB_CONCAT_URL = re.compile(
|
||||
r"(?:(?:Api|api)\.)?getServiceUrl\(\)\s*\+\s*(?P<quote>`)(?P<url>/.*?)(?P=quote)"
|
||||
)
|
||||
MOBILE_HTTP = re.compile(
|
||||
r"http\.(?P<method>Get|Post|Put|Delete|Patch)(?:<[^\n(]*>)?\(\s*"
|
||||
r"(?P<quote>[`'\"])(?P<url>.*?)(?P=quote)"
|
||||
)
|
||||
MOBILE_UNI = re.compile(
|
||||
r"uni\.request\(\s*\{(?:(?!\}\s*\)).)*?\burl\s*:\s*"
|
||||
r"(?P<quote>[`'\"])(?P<url>.*?)(?P=quote)\s*,\s*"
|
||||
r"method\s*:\s*['\"](?P<method>[A-Za-z]+)['\"]",
|
||||
re.DOTALL,
|
||||
)
|
||||
SERVER_CLIENT = re.compile(
|
||||
r"\._execute_async_request\(\s*['\"](?P<method>[A-Za-z]+)['\"]\s*,\s*"
|
||||
r"f?(?P<quote>[`'\"])(?P<url>/.*?)(?P=quote)",
|
||||
re.DOTALL,
|
||||
)
|
||||
SERVER_DIRECT = re.compile(
|
||||
r"f(?P<quote>[`'\"])\{api_url\}(?P<url>/device/address-book/call)(?P=quote)"
|
||||
)
|
||||
JS_EXPRESSION = re.compile(r"\$\{([^{}]+)\}")
|
||||
PYTHON_EXPRESSION = re.compile(r"\{([^{}]+)\}")
|
||||
PATH_PARAMETER = re.compile(r"\{[^/{}]+\}")
|
||||
|
||||
|
||||
class CallSite(NamedTuple):
|
||||
consumer: str
|
||||
method: str
|
||||
path: str
|
||||
source: str
|
||||
|
||||
|
||||
def _source(path: Path, text: str, offset: int) -> str:
|
||||
line = text.count("\n", 0, offset) + 1
|
||||
return f"{path.relative_to(REPO_ROOT).as_posix()}:{line}"
|
||||
|
||||
|
||||
def _parameter_name(expression: str) -> str:
|
||||
identifiers = re.findall(r"[A-Za-z_][A-Za-z0-9_]*", expression)
|
||||
ignored = {"getServiceUrl", "encodeURIComponent", "toString", "value"}
|
||||
useful = [item for item in identifiers if item not in ignored]
|
||||
return useful[-1] if useful else "value"
|
||||
|
||||
|
||||
def normalize_path(raw: str) -> str:
|
||||
value = raw.strip()
|
||||
for marker in (
|
||||
"${getServiceUrl()}",
|
||||
"${Api.getServiceUrl()}",
|
||||
"${api.getServiceUrl()}",
|
||||
"${baseUrlInput.value}",
|
||||
"${getEnvBaseUrl()}",
|
||||
"{api_url}",
|
||||
):
|
||||
if value.startswith(marker):
|
||||
value = value[len(marker) :]
|
||||
break
|
||||
value = value.split("?", 1)[0].split("#", 1)[0]
|
||||
value = JS_EXPRESSION.sub(lambda match: "{" + _parameter_name(match.group(1)) + "}", value)
|
||||
value = PYTHON_EXPRESSION.sub(lambda match: "{" + _parameter_name(match.group(1)) + "}", value)
|
||||
if not value.startswith("/"):
|
||||
raise ValueError(f"consumer URL does not resolve to a manager-api path: {raw!r}")
|
||||
return re.sub(r"/{2,}", "/", value).rstrip("/") or "/"
|
||||
|
||||
|
||||
def _iter_source_files(root: Path, suffixes: set[str]) -> list[Path]:
|
||||
return sorted(
|
||||
path
|
||||
for path in root.rglob("*")
|
||||
if path.is_file()
|
||||
and path.suffix in suffixes
|
||||
and "node_modules" not in path.parts
|
||||
and "dist" not in path.parts
|
||||
)
|
||||
|
||||
|
||||
def extract_web() -> list[CallSite]:
|
||||
calls: list[CallSite] = []
|
||||
for path in _iter_source_files(WEB_ROOT, {".js", ".mjs", ".vue"}):
|
||||
text = path.read_text(encoding="utf-8")
|
||||
covered: list[tuple[int, int]] = []
|
||||
for pattern in (WEB_CHAIN, WEB_CONFIG):
|
||||
for match in pattern.finditer(text):
|
||||
raw_url = match.group("url")
|
||||
if "getServiceUrl()" not in raw_url:
|
||||
continue
|
||||
calls.append(
|
||||
CallSite(
|
||||
"manager-web",
|
||||
match.group("method").upper(),
|
||||
normalize_path(raw_url),
|
||||
_source(path, text, match.start("url")),
|
||||
)
|
||||
)
|
||||
covered.append(match.span("url"))
|
||||
|
||||
def already_covered(offset: int, spans: list[tuple[int, int]] = covered) -> bool:
|
||||
return any(start <= offset < end for start, end in spans)
|
||||
|
||||
for pattern in (WEB_TEMPLATE_URL, WEB_CONCAT_URL):
|
||||
for match in pattern.finditer(text):
|
||||
if already_covered(match.start("url")):
|
||||
continue
|
||||
line_start = text.rfind("\n", 0, match.start()) + 1
|
||||
line_end = text.find("\n", match.end())
|
||||
line = text[line_start : len(text) if line_end < 0 else line_end]
|
||||
if "console.log" in line:
|
||||
continue
|
||||
calls.append(
|
||||
CallSite(
|
||||
"manager-web",
|
||||
"GET",
|
||||
normalize_path(match.group("url")),
|
||||
_source(path, text, match.start("url")),
|
||||
)
|
||||
)
|
||||
|
||||
literal_builders = len(
|
||||
re.findall(r"\.url\(\s*[`'\"]\$\{getServiceUrl\(\)\}", text)
|
||||
)
|
||||
parsed_builders = sum(
|
||||
1
|
||||
for match in WEB_CHAIN.finditer(text)
|
||||
if "getServiceUrl()" in match.group("url")
|
||||
)
|
||||
if literal_builders != parsed_builders:
|
||||
raise RuntimeError(
|
||||
f"unparsed manager-web request builder(s) in {path}: "
|
||||
f"found={literal_builders}, parsed={parsed_builders}"
|
||||
)
|
||||
return calls
|
||||
|
||||
|
||||
def extract_mobile() -> list[CallSite]:
|
||||
calls: list[CallSite] = []
|
||||
for path in _iter_source_files(MOBILE_ROOT, {".ts", ".vue"}):
|
||||
text = path.read_text(encoding="utf-8")
|
||||
parsed_http = list(MOBILE_HTTP.finditer(text))
|
||||
raw_http_count = len(re.findall(r"\bhttp\.(?:Get|Post|Put|Delete|Patch)(?:<|\()", text))
|
||||
if raw_http_count != len(parsed_http):
|
||||
raise RuntimeError(
|
||||
f"unparsed manager-mobile http call(s) in {path}: "
|
||||
f"found={raw_http_count}, parsed={len(parsed_http)}"
|
||||
)
|
||||
for match in parsed_http:
|
||||
calls.append(
|
||||
CallSite(
|
||||
"manager-mobile",
|
||||
match.group("method").upper(),
|
||||
normalize_path(match.group("url")),
|
||||
_source(path, text, match.start("url")),
|
||||
)
|
||||
)
|
||||
for match in MOBILE_UNI.finditer(text):
|
||||
raw_url = match.group("url")
|
||||
if not raw_url.startswith(("${baseUrlInput.value}", "${getEnvBaseUrl()}")):
|
||||
continue
|
||||
calls.append(
|
||||
CallSite(
|
||||
"manager-mobile",
|
||||
match.group("method").upper(),
|
||||
normalize_path(raw_url),
|
||||
_source(path, text, match.start("url")),
|
||||
)
|
||||
)
|
||||
return calls
|
||||
|
||||
|
||||
def extract_server() -> list[CallSite]:
|
||||
calls: list[CallSite] = []
|
||||
path = SERVER_ROOT / "config" / "manage_api_client.py"
|
||||
text = path.read_text(encoding="utf-8")
|
||||
matches = list(SERVER_CLIENT.finditer(text))
|
||||
raw_count = len(re.findall(r"\._execute_async_request\(", text))
|
||||
if raw_count != len(matches):
|
||||
raise RuntimeError(
|
||||
f"unparsed xiaozhi-server manager client call(s): found={raw_count}, parsed={len(matches)}"
|
||||
)
|
||||
for match in matches:
|
||||
calls.append(
|
||||
CallSite(
|
||||
"xiaozhi-server",
|
||||
match.group("method").upper(),
|
||||
normalize_path(match.group("url")),
|
||||
_source(path, text, match.start("url")),
|
||||
)
|
||||
)
|
||||
|
||||
direct_path = SERVER_ROOT / "plugins_func" / "functions" / "call_device.py"
|
||||
direct_text = direct_path.read_text(encoding="utf-8")
|
||||
direct_matches = list(SERVER_DIRECT.finditer(direct_text))
|
||||
if len(direct_matches) != 1:
|
||||
raise RuntimeError(f"expected one direct manager-api call in {direct_path}, found {len(direct_matches)}")
|
||||
match = direct_matches[0]
|
||||
calls.append(
|
||||
CallSite(
|
||||
"xiaozhi-server",
|
||||
"GET",
|
||||
normalize_path(match.group("url")),
|
||||
_source(direct_path, direct_text, match.start("url")),
|
||||
)
|
||||
)
|
||||
return calls
|
||||
|
||||
|
||||
def canonical_route(method: str, path: str) -> tuple[str, str]:
|
||||
return method, PATH_PARAMETER.sub("{}", path)
|
||||
|
||||
|
||||
def build_manifest() -> dict[str, object]:
|
||||
calls = sorted(
|
||||
extract_web() + extract_mobile() + extract_server(),
|
||||
key=lambda item: (item.consumer, item.source, item.method, item.path),
|
||||
)
|
||||
consumers: dict[str, dict[str, object]] = {}
|
||||
for consumer in ("manager-web", "manager-mobile", "xiaozhi-server"):
|
||||
selected = [item for item in calls if item.consumer == consumer]
|
||||
consumers[consumer] = {
|
||||
"callSites": len(selected),
|
||||
"uniqueRoutes": len({canonical_route(item.method, item.path) for item in selected}),
|
||||
"methods": dict(sorted(Counter(item.method for item in selected).items())),
|
||||
}
|
||||
all_routes = {canonical_route(item.method, item.path) for item in calls}
|
||||
return {
|
||||
"count": len(calls),
|
||||
"uniqueRoutes": len(all_routes),
|
||||
"consumers": consumers,
|
||||
"calls": [item._asdict() for item in calls],
|
||||
}
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
print(json.dumps(build_manifest(), ensure_ascii=False, indent=2, sort_keys=False))
|
||||
@@ -0,0 +1,168 @@
|
||||
#!/usr/bin/env python3
|
||||
"""Extract the Spring MVC contract without starting the Java application."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import argparse
|
||||
import fnmatch
|
||||
import json
|
||||
import re
|
||||
from dataclasses import asdict, dataclass
|
||||
from pathlib import Path
|
||||
|
||||
MAPPING_RE = re.compile(r"@(Get|Post|Put|Delete|Patch)Mapping(?:\((.*)\))?")
|
||||
CLASS_MAPPING_RE = re.compile(r"@RequestMapping\(\s*(?:value\s*=\s*)?[\"']([^\"']*)[\"']")
|
||||
PATH_RE = re.compile(r"[\"']([^\"']*)[\"']")
|
||||
METHOD_RE = re.compile(r"\bpublic\s+(?:<[^>]+>\s+)?[^=(;]+?\s+(\w+)\s*\(")
|
||||
PERMISSION_RE = re.compile(r'@RequiresPermissions\(\s*"([^"]+)"')
|
||||
|
||||
PUBLIC_PATTERNS = (
|
||||
"/ota/**",
|
||||
"/otaMag/download/**",
|
||||
"/webjars/**",
|
||||
"/druid/**",
|
||||
"/v3/api-docs/**",
|
||||
"/doc.html",
|
||||
"/favicon.ico",
|
||||
"/user/captcha",
|
||||
"/user/smsVerification",
|
||||
"/user/login",
|
||||
"/user/pub-config",
|
||||
"/user/register",
|
||||
"/user/retrieve-password",
|
||||
"/agent/chat-history/download/**",
|
||||
"/agent/play/**",
|
||||
"/voiceClone/play/**",
|
||||
)
|
||||
SERVER_PATTERNS = (
|
||||
"/config/**",
|
||||
"/device/address-book/call",
|
||||
"/agent/chat-history/report",
|
||||
"/agent/chat-summary/**",
|
||||
"/agent/chat-title/**",
|
||||
)
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class Route:
|
||||
method: str
|
||||
path: str
|
||||
controller: str
|
||||
handler: str
|
||||
auth: str
|
||||
permission: str | None
|
||||
source: str
|
||||
line: int
|
||||
java_signature: str
|
||||
|
||||
|
||||
def _spring_match(path: str, pattern: str) -> bool:
|
||||
return fnmatch.fnmatchcase(path, pattern.replace("**", "*"))
|
||||
|
||||
|
||||
def classify_auth(path: str) -> str:
|
||||
if any(_spring_match(path, pattern) for pattern in PUBLIC_PATTERNS):
|
||||
return "anonymous"
|
||||
if any(_spring_match(path, pattern) for pattern in SERVER_PATTERNS):
|
||||
return "server-secret"
|
||||
return "database-token"
|
||||
|
||||
|
||||
def join_paths(base: str, child: str) -> str:
|
||||
if not base:
|
||||
base = "/"
|
||||
if not base.startswith("/"):
|
||||
base = "/" + base
|
||||
if not child:
|
||||
return base
|
||||
if not child.startswith("/"):
|
||||
child = "/" + child
|
||||
if base == "/":
|
||||
return child
|
||||
return base.rstrip("/") + child
|
||||
|
||||
|
||||
def extract_controller(path: Path, root: Path) -> list[Route]:
|
||||
source = path.read_text(encoding="utf-8")
|
||||
class_position = source.find(" class ")
|
||||
class_header = source[:class_position] if class_position >= 0 else source
|
||||
class_match = list(CLASS_MAPPING_RE.finditer(class_header))
|
||||
base = class_match[-1].group(1) if class_match else ""
|
||||
controller_match = re.search(r"public\s+class\s+(\w+)", source)
|
||||
controller = controller_match.group(1) if controller_match else path.stem
|
||||
lines = source.splitlines()
|
||||
routes: list[Route] = []
|
||||
for index, line in enumerate(lines):
|
||||
mapping = MAPPING_RE.search(line)
|
||||
if not mapping:
|
||||
continue
|
||||
method = mapping.group(1).upper()
|
||||
arguments = mapping.group(2) or ""
|
||||
path_match = PATH_RE.search(arguments)
|
||||
child = path_match.group(1) if path_match else ""
|
||||
decorator_block: list[str] = [line]
|
||||
signature_lines: list[str] = []
|
||||
handler = "unknown"
|
||||
for cursor in range(index + 1, min(index + 40, len(lines))):
|
||||
candidate = lines[cursor]
|
||||
if not signature_lines and candidate.lstrip().startswith("@"):
|
||||
decorator_block.append(candidate)
|
||||
continue
|
||||
signature_lines.append(candidate.strip())
|
||||
signature = " ".join(signature_lines)
|
||||
handler_match = METHOD_RE.search(signature)
|
||||
if handler_match:
|
||||
handler = handler_match.group(1)
|
||||
break
|
||||
if "{" in candidate and "public " not in signature:
|
||||
break
|
||||
permission_match = PERMISSION_RE.search("\n".join(decorator_block))
|
||||
route_path = join_paths(base, child)
|
||||
routes.append(
|
||||
Route(
|
||||
method=method,
|
||||
path=route_path,
|
||||
controller=controller,
|
||||
handler=handler,
|
||||
auth=classify_auth(route_path),
|
||||
permission=permission_match.group(1) if permission_match else None,
|
||||
source=str(path.relative_to(root)),
|
||||
line=index + 1,
|
||||
java_signature=" ".join(signature_lines),
|
||||
)
|
||||
)
|
||||
return routes
|
||||
|
||||
|
||||
def extract_routes(java_root: Path, repository_root: Path) -> list[Route]:
|
||||
routes: list[Route] = []
|
||||
for controller in sorted(java_root.rglob("*Controller.java")):
|
||||
routes.extend(extract_controller(controller, repository_root))
|
||||
return sorted(routes, key=lambda route: (route.path, route.method, route.controller, route.handler))
|
||||
|
||||
|
||||
def main() -> int:
|
||||
parser = argparse.ArgumentParser()
|
||||
parser.add_argument("--repository-root", type=Path, default=Path(__file__).resolve().parents[3])
|
||||
parser.add_argument("--output", type=Path)
|
||||
args = parser.parse_args()
|
||||
repository_root = args.repository_root.resolve()
|
||||
java_root = repository_root / "main" / "manager-api" / "src" / "main" / "java"
|
||||
routes = extract_routes(java_root, repository_root)
|
||||
payload = {
|
||||
"source": "main/manager-api",
|
||||
"contextPath": "/xiaozhi",
|
||||
"count": len(routes),
|
||||
"routes": [asdict(route) for route in routes],
|
||||
}
|
||||
serialized = json.dumps(payload, ensure_ascii=False, indent=2) + "\n"
|
||||
if args.output:
|
||||
args.output.parent.mkdir(parents=True, exist_ok=True)
|
||||
args.output.write_text(serialized, encoding="utf-8")
|
||||
else:
|
||||
print(serialized, end="")
|
||||
return 0
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
raise SystemExit(main())
|
||||
+146
@@ -0,0 +1,146 @@
|
||||
#!/bin/sh
|
||||
set -eu
|
||||
|
||||
SCRIPT_DIR=$(CDPATH= cd -- "$(dirname -- "$0")" && pwd)
|
||||
TARGET_DIR=$(CDPATH= cd -- "${SCRIPT_DIR}/.." && pwd)
|
||||
REPOSITORY_ROOT=$(CDPATH= cd -- "${TARGET_DIR}/../.." && pwd)
|
||||
RUNTIME="${REPOSITORY_ROOT}/.runtime"
|
||||
STATE_DIR="${TARGET_DIR}/.test-runtime"
|
||||
MYSQL_BASE="${RUNTIME}/mysql"
|
||||
MYSQL_PORT="${TEST_MYSQL_PORT:-13316}"
|
||||
MYSQL_DATA="${STATE_DIR}/mysql-data"
|
||||
MYSQL_SOCKET="${STATE_DIR}/mysql.sock"
|
||||
MYSQL_PID="${STATE_DIR}/mysql.pid"
|
||||
MYSQL_LOG="${STATE_DIR}/mysql.log"
|
||||
REDIS_PORT="${TEST_REDIS_PORT:-16379}"
|
||||
REDIS_DIR="${STATE_DIR}/redis-data"
|
||||
REDIS_PID="${STATE_DIR}/redis.pid"
|
||||
REDIS_LOG="${STATE_DIR}/redis.log"
|
||||
TEST_USER="xiaozhi_test"
|
||||
TEST_PASSWORD="isolated-test-only"
|
||||
JAVA_DATABASE="manager_java_test"
|
||||
FASTAPI_DATABASE="manager_fastapi_test"
|
||||
|
||||
require_binaries() {
|
||||
for binary in "${MYSQL_BASE}/bin/mysqld" "${MYSQL_BASE}/bin/mysql" \
|
||||
"${RUNTIME}/redis/bin/redis-server" "${RUNTIME}/redis/bin/redis-cli"; do
|
||||
if [ ! -x "${binary}" ]; then
|
||||
echo "Missing isolated-test binary: ${binary}" >&2
|
||||
exit 1
|
||||
fi
|
||||
done
|
||||
}
|
||||
|
||||
mysql_ready() {
|
||||
"${MYSQL_BASE}/bin/mysqladmin" --protocol=SOCKET --socket="${MYSQL_SOCKET}" -uroot ping >/dev/null 2>&1
|
||||
}
|
||||
|
||||
redis_ready() {
|
||||
"${RUNTIME}/redis/bin/redis-cli" -h 127.0.0.1 -p "${REDIS_PORT}" ping >/dev/null 2>&1
|
||||
}
|
||||
|
||||
wait_until() {
|
||||
description=$1
|
||||
shift
|
||||
attempts=0
|
||||
until "$@"; do
|
||||
attempts=$((attempts + 1))
|
||||
if [ "${attempts}" -ge 100 ]; then
|
||||
echo "Timed out waiting for ${description}" >&2
|
||||
exit 1
|
||||
fi
|
||||
sleep 0.1
|
||||
done
|
||||
}
|
||||
|
||||
start_mysql() {
|
||||
mkdir -p "${STATE_DIR}" "${MYSQL_DATA}"
|
||||
if [ ! -d "${MYSQL_DATA}/mysql" ]; then
|
||||
"${MYSQL_BASE}/bin/mysqld" --no-defaults --initialize-insecure \
|
||||
--basedir="${MYSQL_BASE}" --datadir="${MYSQL_DATA}" --log-error="${MYSQL_LOG}"
|
||||
fi
|
||||
if ! mysql_ready; then
|
||||
"${MYSQL_BASE}/bin/mysqld" --no-defaults --daemonize \
|
||||
--basedir="${MYSQL_BASE}" --datadir="${MYSQL_DATA}" --port="${MYSQL_PORT}" \
|
||||
--socket="${MYSQL_SOCKET}" --pid-file="${MYSQL_PID}" --log-error="${MYSQL_LOG}" \
|
||||
--bind-address=127.0.0.1 --mysqlx=0 --skip-name-resolve \
|
||||
--character-set-server=utf8mb4 --collation-server=utf8mb4_unicode_ci \
|
||||
--default-time-zone=+08:00
|
||||
wait_until "isolated MySQL" mysql_ready
|
||||
fi
|
||||
"${MYSQL_BASE}/bin/mysql" --protocol=SOCKET --socket="${MYSQL_SOCKET}" -uroot \
|
||||
-e "CREATE DATABASE IF NOT EXISTS ${JAVA_DATABASE} CHARACTER SET utf8mb4 COLLATE utf8mb4_unicode_ci; CREATE DATABASE IF NOT EXISTS ${FASTAPI_DATABASE} CHARACTER SET utf8mb4 COLLATE utf8mb4_unicode_ci; CREATE USER IF NOT EXISTS '${TEST_USER}'@'127.0.0.1' IDENTIFIED BY '${TEST_PASSWORD}'; GRANT ALL PRIVILEGES ON ${JAVA_DATABASE}.* TO '${TEST_USER}'@'127.0.0.1'; GRANT ALL PRIVILEGES ON ${FASTAPI_DATABASE}.* TO '${TEST_USER}'@'127.0.0.1'; FLUSH PRIVILEGES;"
|
||||
}
|
||||
|
||||
start_redis() {
|
||||
mkdir -p "${REDIS_DIR}"
|
||||
if ! redis_ready; then
|
||||
"${RUNTIME}/redis/bin/redis-server" --daemonize yes --bind 127.0.0.1 \
|
||||
--port "${REDIS_PORT}" --pidfile "${REDIS_PID}" --dir "${REDIS_DIR}" \
|
||||
--logfile "${REDIS_LOG}" --save "" --appendonly no --databases 16
|
||||
wait_until "isolated Redis" redis_ready
|
||||
fi
|
||||
}
|
||||
|
||||
start() {
|
||||
require_binaries
|
||||
start_mysql
|
||||
start_redis
|
||||
echo "Isolated MySQL ${MYSQL_PORT} and Redis ${REDIS_PORT} are ready."
|
||||
}
|
||||
|
||||
stop() {
|
||||
if redis_ready; then
|
||||
"${RUNTIME}/redis/bin/redis-cli" -h 127.0.0.1 -p "${REDIS_PORT}" shutdown nosave >/dev/null
|
||||
fi
|
||||
if mysql_ready; then
|
||||
"${MYSQL_BASE}/bin/mysqladmin" --protocol=SOCKET --socket="${MYSQL_SOCKET}" -uroot shutdown
|
||||
fi
|
||||
echo "Isolated services stopped."
|
||||
}
|
||||
|
||||
reset() {
|
||||
start
|
||||
"${MYSQL_BASE}/bin/mysql" --protocol=SOCKET --socket="${MYSQL_SOCKET}" -uroot \
|
||||
-e "DROP DATABASE IF EXISTS ${JAVA_DATABASE}; DROP DATABASE IF EXISTS ${FASTAPI_DATABASE}; CREATE DATABASE ${JAVA_DATABASE} CHARACTER SET utf8mb4 COLLATE utf8mb4_unicode_ci; CREATE DATABASE ${FASTAPI_DATABASE} CHARACTER SET utf8mb4 COLLATE utf8mb4_unicode_ci; GRANT ALL PRIVILEGES ON ${JAVA_DATABASE}.* TO '${TEST_USER}'@'127.0.0.1'; GRANT ALL PRIVILEGES ON ${FASTAPI_DATABASE}.* TO '${TEST_USER}'@'127.0.0.1';"
|
||||
"${RUNTIME}/redis/bin/redis-cli" -h 127.0.0.1 -p "${REDIS_PORT}" flushall >/dev/null
|
||||
echo "Only the isolated test schemas and isolated Redis instance were reset."
|
||||
}
|
||||
|
||||
migrate_one() {
|
||||
database=$1
|
||||
LIQUIBASE_URL="jdbc:mysql://127.0.0.1:${MYSQL_PORT}/${database}?useUnicode=true&characterEncoding=UTF-8&serverTimezone=Asia/Shanghai&allowMultiQueries=true" \
|
||||
LIQUIBASE_USERNAME="${TEST_USER}" \
|
||||
LIQUIBASE_PASSWORD="${TEST_PASSWORD}" \
|
||||
MAVEN_BIN="${RUNTIME}/maven/bin/mvn" \
|
||||
MAVEN_LOCAL_REPOSITORY="${RUNTIME}/m2" \
|
||||
JAVA_RESOURCES_DIR="${REPOSITORY_ROOT}/main/manager-api/src/main/resources" \
|
||||
"${SCRIPT_DIR}/run-migrations.sh"
|
||||
}
|
||||
|
||||
migrate() {
|
||||
start
|
||||
migrate_one "${JAVA_DATABASE}"
|
||||
migrate_one "${FASTAPI_DATABASE}"
|
||||
}
|
||||
|
||||
print_env() {
|
||||
cat <<EOF
|
||||
export TEST_MYSQL_PORT='${MYSQL_PORT}'
|
||||
export TEST_REDIS_PORT='${REDIS_PORT}'
|
||||
export TEST_JAVA_DATABASE_URL='mysql+asyncmy://${TEST_USER}:${TEST_PASSWORD}@127.0.0.1:${MYSQL_PORT}/${JAVA_DATABASE}?charset=utf8mb4'
|
||||
export TEST_FASTAPI_DATABASE_URL='mysql+asyncmy://${TEST_USER}:${TEST_PASSWORD}@127.0.0.1:${MYSQL_PORT}/${FASTAPI_DATABASE}?charset=utf8mb4'
|
||||
export TEST_JAVA_JDBC_URL='jdbc:mysql://127.0.0.1:${MYSQL_PORT}/${JAVA_DATABASE}?useUnicode=true&characterEncoding=UTF-8&serverTimezone=Asia/Shanghai&allowMultiQueries=true'
|
||||
export TEST_JAVA_REDIS_URL='redis://127.0.0.1:${REDIS_PORT}/1'
|
||||
export TEST_FASTAPI_REDIS_URL='redis://127.0.0.1:${REDIS_PORT}/2'
|
||||
EOF
|
||||
}
|
||||
|
||||
case "${1:-}" in
|
||||
start) start ;;
|
||||
stop) stop ;;
|
||||
reset) reset ;;
|
||||
migrate) migrate ;;
|
||||
env) print_env ;;
|
||||
*) echo "usage: $0 {start|stop|reset|migrate|env}" >&2; exit 2 ;;
|
||||
esac
|
||||
@@ -0,0 +1,640 @@
|
||||
#!/usr/bin/env python3
|
||||
"""Render the auditable Java/FastAPI compatibility matrix from checked-in inventories."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import argparse
|
||||
import json
|
||||
import re
|
||||
from collections import Counter
|
||||
from pathlib import Path
|
||||
from typing import Any
|
||||
|
||||
TARGET_ROOT = Path(__file__).resolve().parents[1]
|
||||
REPOSITORY_ROOT = TARGET_ROOT.parents[1]
|
||||
JAVA_ROOT = REPOSITORY_ROOT / "main" / "manager-api"
|
||||
JAVA_SOURCE_ROOT = JAVA_ROOT / "src" / "main" / "java"
|
||||
JAVA_RESOURCE_ROOT = JAVA_ROOT / "src" / "main" / "resources"
|
||||
JAVA_MANIFEST = TARGET_ROOT / "compatibility" / "java-routes.json"
|
||||
CONSUMER_MANIFEST = TARGET_ROOT / "compatibility" / "consumer-routes.json"
|
||||
CONTRACT_RESULTS = TARGET_ROOT / "compatibility" / "contract-results.json"
|
||||
ROUTE_SURFACE_RESULTS = TARGET_ROOT / "compatibility" / "route-surface-results.json"
|
||||
AUTHENTICATED_ROUTE_RESULTS = TARGET_ROOT / "compatibility" / "authenticated-route-results.json"
|
||||
PATH_PARAMETER = re.compile(r"\{[^/{}]+\}")
|
||||
|
||||
|
||||
DIFFERENTIAL_CASES: dict[tuple[str, str], str] = {
|
||||
("GET", "/user/pub-config"): "1",
|
||||
("GET", "/user/info"): "9(七语言/过期 Token/Long)",
|
||||
("GET", "/admin/users"): "3(权限/序列化/非法分页)",
|
||||
("GET", "/agent/list"): "1",
|
||||
("GET", "/device/bind/{}"): "1",
|
||||
("GET", "/models/provider"): "1",
|
||||
("GET", "/correct-word/file/list"): "1",
|
||||
("GET", "/correct-word/file/download/{}"): "2(二进制/更新后下载)",
|
||||
("PUT", "/device/update/{}"): "3(上下界/UTF-16 长度)",
|
||||
("POST", "/models/provider"): "1(约束集合)",
|
||||
("POST", "/config/server-base"): "3(缺失/错误/正确 secret)",
|
||||
("GET", "/ota/"): "1(MIME/body)",
|
||||
("POST", "/ota/"): "4(必填/格式/凭证/密码学)",
|
||||
("POST", "/ota/activate"): "3",
|
||||
("POST", "/device/tools/list/{}"): "2(响应/外呼格式)",
|
||||
("POST", "/correct-word/file"): "2(响应/DB)",
|
||||
("PUT", "/correct-word/file/{}"): "2(响应/DB)",
|
||||
("DELETE", "/correct-word/file/{}"): "1(级联副作用)",
|
||||
("POST", "/otaMag/upload"): "2(上传/扩展名错误)",
|
||||
("POST", "/otaMag"): "2(响应/DB)",
|
||||
("GET", "/otaMag/download/{}"): "4(次数限制及二进制)",
|
||||
}
|
||||
|
||||
DOMAIN_TESTS = {
|
||||
"AdminController": "sys",
|
||||
"SysParamsController": "sys",
|
||||
"SysDictDataController": "sys",
|
||||
"SysDictTypeController": "sys",
|
||||
"ServerSideManageController": "sys",
|
||||
"LoginController": "security",
|
||||
"ConfigController": "config",
|
||||
"AgentController": "agent",
|
||||
"AgentChatHistoryController": "agent",
|
||||
"AgentMcpAccessPointController": "agent",
|
||||
"AgentSnapshotController": "agent",
|
||||
"AgentTemplateController": "agent",
|
||||
"AgentVoicePrintController": "agent",
|
||||
"CorrectWordController": "correctword",
|
||||
"DeviceController": "device",
|
||||
"KnowledgeBaseController": "knowledge",
|
||||
"KnowledgeFilesController": "knowledge",
|
||||
"ModelController": "model",
|
||||
"ModelProviderController": "model",
|
||||
"OTAController": "device",
|
||||
"OTAMagController": "device",
|
||||
"TimbreController": "timbre",
|
||||
"VoiceCloneController": "voiceclone",
|
||||
"VoiceResourceController": "voiceclone",
|
||||
}
|
||||
|
||||
|
||||
def _canonical(path: str) -> str:
|
||||
return PATH_PARAMETER.sub("{}", path)
|
||||
|
||||
|
||||
def _declaration(route: dict[str, Any]) -> str:
|
||||
source_path = REPOSITORY_ROOT / route["source"]
|
||||
lines = source_path.read_text(encoding="utf-8").splitlines()
|
||||
excerpt = "\n".join(lines[int(route["line"]) - 1 : int(route["line"]) + 60])
|
||||
match = re.search(r"\bpublic\s+", excerpt)
|
||||
if match is None:
|
||||
raise ValueError(f"public declaration not found for {route['controller']}.{route['handler']}")
|
||||
declaration = excerpt[match.start() :]
|
||||
depth = 0
|
||||
saw_parenthesis = False
|
||||
for offset, character in enumerate(declaration):
|
||||
if character == "(":
|
||||
depth += 1
|
||||
saw_parenthesis = True
|
||||
elif character == ")":
|
||||
depth -= 1
|
||||
elif character == "{" and saw_parenthesis and depth == 0:
|
||||
return " ".join(declaration[:offset].split())
|
||||
raise ValueError(f"unterminated declaration for {route['controller']}.{route['handler']}")
|
||||
|
||||
|
||||
def _parameter_text(declaration: str, handler: str) -> str:
|
||||
marker = re.search(rf"\b{re.escape(handler)}\s*\(", declaration)
|
||||
if marker is None:
|
||||
return ""
|
||||
start = marker.end() - 1
|
||||
depth = 0
|
||||
for offset in range(start, len(declaration)):
|
||||
character = declaration[offset]
|
||||
if character == "(":
|
||||
depth += 1
|
||||
elif character == ")":
|
||||
depth -= 1
|
||||
if depth == 0:
|
||||
return declaration[start + 1 : offset]
|
||||
return ""
|
||||
|
||||
|
||||
def _split_parameters(value: str) -> list[str]:
|
||||
result: list[str] = []
|
||||
start = 0
|
||||
round_depth = 0
|
||||
angle_depth = 0
|
||||
in_quote: str | None = None
|
||||
escaped = False
|
||||
for offset, character in enumerate(value):
|
||||
if escaped:
|
||||
escaped = False
|
||||
continue
|
||||
if character == "\\":
|
||||
escaped = True
|
||||
continue
|
||||
if in_quote is not None:
|
||||
if character == in_quote:
|
||||
in_quote = None
|
||||
continue
|
||||
if character in {'"', "'"}:
|
||||
in_quote = character
|
||||
elif character == "(":
|
||||
round_depth += 1
|
||||
elif character == ")":
|
||||
round_depth -= 1
|
||||
elif character == "<":
|
||||
angle_depth += 1
|
||||
elif character == ">":
|
||||
angle_depth -= 1
|
||||
elif character == "," and round_depth == 0 and angle_depth == 0:
|
||||
result.append(value[start:offset].strip())
|
||||
start = offset + 1
|
||||
tail = value[start:].strip()
|
||||
if tail:
|
||||
result.append(tail)
|
||||
return result
|
||||
|
||||
|
||||
def _without_annotations(value: str) -> str:
|
||||
output: list[str] = []
|
||||
offset = 0
|
||||
while offset < len(value):
|
||||
if value[offset] != "@":
|
||||
output.append(value[offset])
|
||||
offset += 1
|
||||
continue
|
||||
offset += 1
|
||||
while offset < len(value) and (value[offset].isalnum() or value[offset] in "._$"):
|
||||
offset += 1
|
||||
while offset < len(value) and value[offset].isspace():
|
||||
offset += 1
|
||||
if offset < len(value) and value[offset] == "(":
|
||||
depth = 1
|
||||
offset += 1
|
||||
in_quote: str | None = None
|
||||
while offset < len(value) and depth:
|
||||
character = value[offset]
|
||||
if in_quote is not None:
|
||||
if character == in_quote and value[offset - 1] != "\\":
|
||||
in_quote = None
|
||||
elif character in {'"', "'"}:
|
||||
in_quote = character
|
||||
elif character == "(":
|
||||
depth += 1
|
||||
elif character == ")":
|
||||
depth -= 1
|
||||
offset += 1
|
||||
while offset < len(value) and value[offset].isspace():
|
||||
offset += 1
|
||||
return " ".join("".join(output).split())
|
||||
|
||||
|
||||
def _type_and_name(parameter: str) -> tuple[str, str]:
|
||||
cleaned = _without_annotations(parameter).removeprefix("final ").strip()
|
||||
pieces = cleaned.rsplit(" ", 1)
|
||||
if len(pieces) != 2:
|
||||
return cleaned, cleaned
|
||||
return pieces[0], pieces[1]
|
||||
|
||||
|
||||
def _request_surface(route: dict[str, Any], declaration: str) -> str:
|
||||
path_names = re.findall(r"\{([^/{}]+)\}", route["path"])
|
||||
headers: list[str] = []
|
||||
queries: list[str] = []
|
||||
bodies: list[str] = []
|
||||
multipart: list[str] = []
|
||||
for parameter in _split_parameters(_parameter_text(declaration, route["handler"])):
|
||||
parameter_type, name = _type_and_name(parameter)
|
||||
if "HttpServletResponse" in parameter_type or "HttpServletRequest" in parameter_type:
|
||||
continue
|
||||
if "@PathVariable" in parameter:
|
||||
continue
|
||||
if "@RequestHeader" in parameter:
|
||||
quoted = re.search(r'@RequestHeader(?:\([^)]*)?["\']([^"\']+)["\']', parameter)
|
||||
headers.append(quoted.group(1) if quoted else name)
|
||||
elif "MultipartFile" in parameter_type:
|
||||
multipart.append(name)
|
||||
elif "@RequestBody" in parameter:
|
||||
bodies.append(parameter_type)
|
||||
elif "@RequestParam" in parameter or "@ParameterObject" in parameter:
|
||||
queries.append(f"{name}:{parameter_type}" if parameter_type != name else name)
|
||||
elif route["method"] == "GET" or route["controller"] == "ModelProviderController":
|
||||
queries.append(f"{name}:{parameter_type}" if parameter_type != name else name)
|
||||
parts: list[str] = []
|
||||
if path_names:
|
||||
parts.append("Path:" + ",".join(path_names))
|
||||
if headers:
|
||||
parts.append("Header:" + ",".join(headers))
|
||||
if queries:
|
||||
parts.append("Query:" + ",".join(queries))
|
||||
if bodies:
|
||||
parts.append("Body:" + ",".join(bodies))
|
||||
if multipart:
|
||||
parts.append("Multipart:" + ",".join(multipart))
|
||||
return "; ".join(parts) if parts else "—"
|
||||
|
||||
|
||||
def _response_type(route: dict[str, Any], declaration: str) -> str:
|
||||
path = route["path"]
|
||||
if path == "/user/captcha":
|
||||
return "image/gif 二进制"
|
||||
if path == "/ota/" and route["method"] == "GET":
|
||||
return "裸 text/plain"
|
||||
if path.startswith("/ota/"):
|
||||
return "裸 application/json"
|
||||
if path in {
|
||||
"/agent/play/{uuid}",
|
||||
"/agent/chat-history/download/{uuid}/current",
|
||||
"/agent/chat-history/download/{uuid}/previous",
|
||||
"/correct-word/file/download/{fileId}",
|
||||
"/otaMag/download/{uuid}",
|
||||
"/voiceClone/play/{uuid}",
|
||||
}:
|
||||
return "流式/二进制 + 原下载 headers"
|
||||
match = re.search(rf"public\s+(.+?)\s+{re.escape(route['handler'])}\s*\(", declaration)
|
||||
return_type = match.group(1) if match else "unknown"
|
||||
if return_type.startswith("Result<"):
|
||||
return "envelope " + return_type.removeprefix("Result")
|
||||
return return_type
|
||||
|
||||
|
||||
def _permission(route: dict[str, Any]) -> str | None:
|
||||
source_path = REPOSITORY_ROOT / route["source"]
|
||||
lines = source_path.read_text(encoding="utf-8").splitlines()
|
||||
excerpt = "\n".join(lines[int(route["line"]) - 1 : int(route["line"]) + 60])
|
||||
declaration_offset = excerpt.find("public ")
|
||||
decorators = excerpt if declaration_offset < 0 else excerpt[:declaration_offset]
|
||||
match = re.search(r'@RequiresPermissions\(\s*"([^"]+)"', decorators)
|
||||
return match.group(1) if match else None
|
||||
|
||||
|
||||
def _side_effect(route: dict[str, Any]) -> str:
|
||||
method = route["method"]
|
||||
path = route["path"]
|
||||
handler = route["handler"]
|
||||
controller = route["controller"]
|
||||
if path == "/user/captcha":
|
||||
return "Redis-W(captcha TTL); GIF"
|
||||
if path == "/user/login":
|
||||
return "DB-R/W(token); Redis-R/DEL(captcha)"
|
||||
if path == "/user/smsVerification":
|
||||
return "Redis-R/W(TTL/频控); 外部-Aliyun SMS"
|
||||
if path in {"/user/register", "/user/retrieve-password", "/user/change-password"}:
|
||||
return "DB-W(user/token); Redis-R/DEL(SMS)"
|
||||
if path in {"/user/info", "/user/pub-config"}:
|
||||
return "DB-R; Redis-R/W(cache)"
|
||||
if path.startswith("/admin/server/"):
|
||||
return "DB/Redis-R(secret/WS); Redis-W(one-shot); 外部-WebSocket" if method == "POST" else "DB/Redis-R"
|
||||
if path.startswith("/admin/params"):
|
||||
if method == "GET":
|
||||
return "DB-R"
|
||||
if method == "PUT":
|
||||
return "DB-W; Redis-W; 外部-配置端点探测(按 paramCode)"
|
||||
return "DB-W; Redis-W/DEL"
|
||||
if path.startswith("/admin/dict"):
|
||||
return "DB-R; Redis-R/W(dict cache)" if method == "GET" else "DB-W; Redis-DEL(dict cache)"
|
||||
if path == "/admin/device/all" or path == "/admin/users":
|
||||
return "DB-R"
|
||||
if path == "/admin/users/{id}" and method == "DELETE":
|
||||
return "DB-W(用户/token/device/agent 级联)"
|
||||
if path.startswith("/admin/users"):
|
||||
return "DB-W(user/password/status/token)"
|
||||
if path.startswith("/config/"):
|
||||
return "DB-R; Redis-R/W(runtime/model/timbre cache)"
|
||||
if path.startswith("/agent/mcp/tools"):
|
||||
return "DB/Redis-R; 外部-WebSocket MCP"
|
||||
if path.startswith("/agent/mcp/address"):
|
||||
return "DB/Redis-R; AES token 生成"
|
||||
if path.startswith("/agent/voice-print"):
|
||||
return "DB-R" if method == "GET" else "DB-W; 外部-voiceprint HTTP"
|
||||
if "/chat-summary/" in path or "/chat-title/" in path:
|
||||
return "DB-R/W(chat); 外部-OpenAI-compatible LLM"
|
||||
if path == "/agent/chat-history/report":
|
||||
return "DB-W(chat/session); server-secret"
|
||||
if path.startswith("/agent/chat-history/getDownloadUrl/"):
|
||||
return "DB-R(chat/session); Redis-W(download token TTL)"
|
||||
if "/chat-history/download/" in path or path.startswith("/agent/play/"):
|
||||
return "DB/Redis-R(one-shot); 文件-R/流式"
|
||||
if path.startswith("/agent/audio/"):
|
||||
return "DB-R(audio); Redis-W(one-shot URL)"
|
||||
if path.startswith("/agent/") or path == "/agent":
|
||||
return "DB-R; Redis-R" if method == "GET" else "DB-W(含快照/映射/标签事务); Redis-DEL"
|
||||
if path.startswith("/correct-word/"):
|
||||
if "download" in path:
|
||||
return "DB-R(content); 二进制"
|
||||
return "DB-R" if method == "GET" else "DB-W(file/items/mapping 事务)"
|
||||
if path.startswith("/datasets"):
|
||||
if path == "/datasets/rag-models":
|
||||
return "DB-R(model config)"
|
||||
if method == "GET" and not path.endswith("/chunks"):
|
||||
return "DB-R"
|
||||
return "DB-R/W; 外部-RAGFlow HTTP(upload/dataset/document/chunk/retrieval)"
|
||||
if path.startswith("/device/tools/") or (path == "/device/bind/{agentId}" and method == "POST"):
|
||||
return "DB/Redis-R; 外部-MQTT gateway HTTP + daily auth"
|
||||
if path in {"/device/address-book/call", "/device/address-book/lookup"}:
|
||||
return "DB-R; 外部-MQTT gateway HTTP; server-secret"
|
||||
if path.startswith("/device/"):
|
||||
return "DB-R; Redis-R" if method == "GET" else "DB-W(device/bind/address-book); Redis-R/W"
|
||||
if path == "/ota/":
|
||||
if method == "GET":
|
||||
return "—"
|
||||
return "DB/Redis-R(设备/固件/配置); HMAC/Base64/时间戳凭证"
|
||||
if path == "/ota/activate":
|
||||
return "DB-R/W(device activation); Redis-R/W(TTL)"
|
||||
if path.startswith("/otaMag/upload"):
|
||||
return "文件-W(MD5/扩展名/大小)"
|
||||
if path.startswith("/otaMag/download"):
|
||||
return "Redis-R/W(一次性/次数); 文件-R/流式"
|
||||
if path.startswith("/otaMag/getDownloadUrl"):
|
||||
return "DB-R; Redis-W(download token TTL)"
|
||||
if path.startswith("/otaMag"):
|
||||
if method == "GET":
|
||||
return "DB-R"
|
||||
if method == "DELETE":
|
||||
return "DB-W(OTA metadata); 文件-DEL"
|
||||
return "DB-W(OTA metadata)"
|
||||
if path.startswith("/models"):
|
||||
return "DB-R; Redis-R/W(model cache)" if method == "GET" else "DB-W; Redis-DEL(model/config cache)"
|
||||
if path.startswith("/ttsVoice"):
|
||||
return "DB-R; Redis-R/W(timbre cache)" if method == "GET" else "DB-W; Redis-DEL(timbre/config cache)"
|
||||
if path.startswith("/voiceClone"):
|
||||
if path.startswith("/voiceClone/play"):
|
||||
return "Redis-R/DEL(one-shot); 文件/外部音频-R"
|
||||
if handler == "getAudioId":
|
||||
return "DB-R; Redis-W(one-shot URL)"
|
||||
if handler == "updateName":
|
||||
return "DB-W(train record name)"
|
||||
if method == "GET":
|
||||
return "DB-R"
|
||||
return "DB-R/W(train state); 文件-W; 外部-火山语音克隆 HTTP"
|
||||
if path.startswith("/voiceResource"):
|
||||
return "DB-R" if method == "GET" else "DB-W(voice resource)"
|
||||
operation = "DB-R" if method == "GET" else "DB-W"
|
||||
return f"{operation} ({controller}.{handler})"
|
||||
|
||||
|
||||
def _verification(route: dict[str, Any]) -> str:
|
||||
domain = DOMAIN_TESTS.get(route["controller"])
|
||||
domain_status = f"领域✓({domain},域级)" if domain else "领域—"
|
||||
diff = DIFFERENTIAL_CASES.get((route["method"], _canonical(route["path"])))
|
||||
diff_status = f"差分✓{diff}" if diff else "差分—"
|
||||
if route["path"] == "/otaMag/getDownloadUrl/{id}":
|
||||
diff_status = "差分间接✓(供下载链路)"
|
||||
return f"结构✓;请求面差分✓1;认证业务面差分✓1;{domain_status};{diff_status}"
|
||||
|
||||
|
||||
def _escape(value: str) -> str:
|
||||
return value.replace("|", "\\|").replace("\n", " ")
|
||||
|
||||
|
||||
def _inventory_section(java_routes: list[dict[str, Any]]) -> str:
|
||||
controller_counts = Counter(route["controller"] for route in java_routes)
|
||||
mapper_files = sorted((JAVA_RESOURCE_ROOT / "mapper").rglob("*.xml"))
|
||||
changelog_root = JAVA_RESOURCE_ROOT / "db" / "changelog"
|
||||
sql_files = sorted(changelog_root.rglob("*.sql"))
|
||||
master = changelog_root / "db.changelog-master.yaml"
|
||||
changeset_refs = master.read_text(encoding="utf-8").count("changeSet:")
|
||||
entity_files = list(JAVA_SOURCE_ROOT.rglob("entity/*.java"))
|
||||
dto_files = [path for path in JAVA_SOURCE_ROOT.rglob("*.java") if "dto" in path.parts]
|
||||
vo_files = list(JAVA_SOURCE_ROOT.rglob("vo/*.java"))
|
||||
dao_files = list(JAVA_SOURCE_ROOT.rglob("dao/*.java"))
|
||||
service_files = list(JAVA_SOURCE_ROOT.rglob("service/**/*.java"))
|
||||
service_impl_files = list(JAVA_SOURCE_ROOT.rglob("service/impl/*.java"))
|
||||
assert len(controller_counts) == 24
|
||||
assert len(mapper_files) == 20
|
||||
assert len(sql_files) == 101
|
||||
assert changeset_refs == 101
|
||||
assert len(entity_files) == 29
|
||||
assert len(dto_files) == 58
|
||||
assert len(vo_files) == 14
|
||||
assert len(dao_files) == 29
|
||||
mapper_names = "、".join(f"`{path.relative_to(JAVA_RESOURCE_ROOT).as_posix()}`" for path in mapper_files)
|
||||
controller_names = "、".join(
|
||||
f"`{controller}`({controller_counts[controller]})" for controller in sorted(controller_counts)
|
||||
)
|
||||
return f"""## Java 基线静态盘点
|
||||
|
||||
- Controller:24 个、154 条映射。按 Controller 的路由数为:{controller_names}。
|
||||
- 数据分层:`entity/` 29 个 Java 文件(28 个 `*Entity.java` 加 `BaseEntity`)、`dto/` 58 个、
|
||||
`vo/` 14 个、`dao/` 29 个、`service/` 树 {len(service_files)} 个文件(其中
|
||||
`service/impl/` {len(service_impl_files)} 个)。FastAPI 对应落在 `schemas/`、`repositories/`、
|
||||
`services/`、`routers/`、`integrations/` 与 `jobs/`,没有把跨表事务放进路由。
|
||||
- MyBatis XML:20 个,分别是 {mapper_names}。
|
||||
- Liquibase:`db.changelog-master.yaml` 含 {changeset_refs} 个 `changeSet` 引用,目录中恰有
|
||||
{len(sql_files)} 个 SQL;Python 部署继续执行这 101 个原始 SQL,不改写历史。
|
||||
- 定时工作:`DocumentStatusSyncTask` 每次完成后延迟 30 秒,扫描 RAGFlow RUNNING 文档并
|
||||
回写 SUCCESS/FAIL/CANCEL 与统计;当前 Java 源码另有 `AgentSnapshotRedactionRunner`,启动时
|
||||
执行一次并在滚动部署期每 15 秒补偿脱敏旧快照。FastAPI 将工作移到独立 jobs 进程,并以
|
||||
Redis 分布式锁/watchdog 防止多 worker 重复执行。
|
||||
- 外部集成:RAGFlow dataset/document/chunk/retrieval/upload;阿里云短信;火山语音克隆训练与
|
||||
音频;声纹 HTTP;OpenAI-compatible LLM 摘要/标题;MQTT gateway HTTP;MCP/管理动作
|
||||
WebSocket;OTA/WS/MQTT 的 HMAC、Base64、时间戳与下载文件存储。自动测试只访问可重复 mock,
|
||||
未使用真实付费凭证。
|
||||
"""
|
||||
|
||||
|
||||
def _require_complete_route_report(
|
||||
report: dict[str, Any],
|
||||
routes: list[dict[str, Any]],
|
||||
*,
|
||||
request_profile: str,
|
||||
side_effect_policy: str,
|
||||
) -> None:
|
||||
expected_summary = {"total": 154, "passed": 154, "failed": 0, "skipped": 0}
|
||||
expected_names = [f"{route['method']} {route['path']}" for route in routes]
|
||||
assert report["summary"] == expected_summary
|
||||
assert report["coverage"] == {
|
||||
"java_routes": 154,
|
||||
"request_profile": request_profile,
|
||||
"side_effect_policy": side_effect_policy,
|
||||
}
|
||||
assert [result["name"] for result in report["results"]] == expected_names
|
||||
assert all(result["passed"] is True and result["difference"] is None for result in report["results"])
|
||||
|
||||
|
||||
def render() -> str:
|
||||
java_manifest = json.loads(JAVA_MANIFEST.read_text(encoding="utf-8"))
|
||||
consumer_manifest = json.loads(CONSUMER_MANIFEST.read_text(encoding="utf-8"))
|
||||
contract = json.loads(CONTRACT_RESULTS.read_text(encoding="utf-8"))
|
||||
route_surface = json.loads(ROUTE_SURFACE_RESULTS.read_text(encoding="utf-8"))
|
||||
authenticated_route = json.loads(AUTHENTICATED_ROUTE_RESULTS.read_text(encoding="utf-8"))
|
||||
routes: list[dict[str, Any]] = java_manifest["routes"]
|
||||
assert len(routes) == 154
|
||||
assert contract["summary"] == {"total": 49, "passed": 49, "failed": 0, "skipped": 0}
|
||||
_require_complete_route_report(
|
||||
route_surface,
|
||||
routes,
|
||||
request_profile="missing-auth-or-safe-invalid-input",
|
||||
side_effect_policy="no successful write request is issued",
|
||||
)
|
||||
_require_complete_route_report(
|
||||
authenticated_route,
|
||||
routes,
|
||||
request_profile="authenticated-safe-business-or-validation",
|
||||
side_effect_policy="no intentional successful writes",
|
||||
)
|
||||
assert len(DIFFERENTIAL_CASES) == 21
|
||||
|
||||
lines = [
|
||||
"# manager-api FastAPI 兼容性矩阵",
|
||||
"",
|
||||
"> 生成依据:`main/manager-api-fastapi/compatibility/java-routes.json`、",
|
||||
"> `main/manager-api-fastapi/compatibility/consumer-routes.json`、`route-surface-results.json`、",
|
||||
"> `authenticated-route-results.json`、`contract-results.json` 和当前 Java 源码。接口路径均省略",
|
||||
"> 共同前缀 `/xiaozhi`。",
|
||||
"",
|
||||
"## 结论与状态口径",
|
||||
"",
|
||||
"Java 基线共有 **154** 条 Spring MVC 路由;FastAPI 已注册 **154/154(100%)**,并由",
|
||||
"`tests/test_java_route_manifest.py` 对源码清单 freshness、数量和 method/path 注册闭合进行检查。",
|
||||
"此外实现 3 条仅由仓库消费者使用、Java Controller 中不存在的兼容路由,因此这 3 条不计入",
|
||||
"154 条 Java 覆盖率。三端 188 个调用点均能解析到 FastAPI 路由。",
|
||||
"",
|
||||
"矩阵状态必须按下列含义阅读:",
|
||||
"",
|
||||
"- `结构✓`:method/path 已注册且清单闭合;它不等于业务行为逐接口实测。",
|
||||
"- `请求面差分✓1`:本行已向隔离 Java/FastAPI 各发送一次缺少鉴权或安全非法输入,精确比较",
|
||||
" HTTP status、body 与 Content-Type;最终为 **154/154 通过、0 失败、0 跳过**,且不发送成功写请求。",
|
||||
"- `认证业务面差分✓1`:本行已使用有效 DB Token、server-secret 或匿名身份,再向隔离",
|
||||
" Java/FastAPI 各发送一次安全业务/校验请求,精确比较 HTTP status、body 与 Content-Type;",
|
||||
" 最终为 **154/154 通过、0 失败、0 跳过**,且不主动发送成功写请求。该状态不等于每条路由的",
|
||||
" 完整成功生命周期均已差分,完整副作用证据仍以 `差分✓N` 为准。",
|
||||
"- `领域✓(x,域级)`:该领域有 service/repository/协议自动测试,但不保证本行每条成功与错误路径",
|
||||
" 都被直接请求。`领域—` 表示除结构测试外没有可归属的域级直接测试证据。",
|
||||
"- `差分✓N`:本行除安全请求面外,还参与了成功、主要错误、协议或数据库副作用的深度对照;",
|
||||
" 括号说明覆盖面。深度结果为 **49/49 checks 通过、0 失败、0 跳过**,覆盖 **21/154** 条路由。",
|
||||
" `差分间接✓` 表示 J125 作为下载链路的 URL 生成步骤被间接覆盖;`差分—` 表示没有深度对照,",
|
||||
" 不能把 154/154 请求面差分误读成 154 条全部成功路径与副作用都已逐接口对照。",
|
||||
"- 所有 `Result<T>` 均表示 `{code,msg,data}` envelope;原 Java 为 HTTP 200 的认证、权限、业务和",
|
||||
" 参数错误由全局兼容层维持 HTTP 200。二进制/OTA 裸响应在“响应类型”列单独标明。",
|
||||
"",
|
||||
"## 三端消费者闭合",
|
||||
"",
|
||||
"| 消费者 | 调用点 | 唯一结构路由 | 方法分布 |",
|
||||
"|---|---:|---:|---|",
|
||||
]
|
||||
for consumer in ("manager-web", "manager-mobile", "xiaozhi-server"):
|
||||
item = consumer_manifest["consumers"][consumer]
|
||||
methods = "、".join(f"{method} {count}" for method, count in item["methods"].items())
|
||||
lines.append(f"| `{consumer}` | {item['callSites']} | {item['uniqueRoutes']} | {methods} |")
|
||||
lines.extend(
|
||||
[
|
||||
f"| **合计** | **{consumer_manifest['count']}** | **{consumer_manifest['uniqueRoutes']}** | — |",
|
||||
"",
|
||||
"### 3 条消费者孤儿兼容路由",
|
||||
"",
|
||||
"| Method/path | 来源 | FastAPI 语义 | 鉴权 | 状态 |",
|
||||
"|---|---|---|---|---|",
|
||||
(
|
||||
"| `GET /api/ping` | manager-mobile 环境设置探活 | "
|
||||
"`{code:0,msg:\"success\",data:\"pong\"}` | 匿名 | 实现✓;consumer resolve✓ |"
|
||||
),
|
||||
(
|
||||
"| `PUT /user/configDevice/{device_id}` | manager-web 遗留设备配置调用 | "
|
||||
"按现有设备更新契约处理 body | DB Token | 实现✓;consumer resolve✓ |"
|
||||
),
|
||||
(
|
||||
"| `GET /device/address-book/lookup` | xiaozhi-server 管理客户端 | "
|
||||
"`callerMac/nickname/answer` 地址簿查询/呼叫兼容别名 | server-secret | "
|
||||
"实现✓;consumer resolve✓;device 域测试✓ |"
|
||||
),
|
||||
"",
|
||||
"`GET /admin/dict/data/type/FIRMWARE_TYPE` 是动态 Java 路由",
|
||||
"`GET /admin/dict/data/type/{dictType}` 的一个字面调用,不是第四条孤儿路由。",
|
||||
"",
|
||||
_inventory_section(routes).rstrip(),
|
||||
"",
|
||||
"## 154 条 Java→FastAPI 逐接口矩阵",
|
||||
"",
|
||||
"副作用缩写:`DB-R/W`=数据库读/写,`Redis-R/W/DEL`=缓存读/写/失效,`文件-R/W`=文件",
|
||||
"读取/写入;外部调用均在 service/integration 层。权限为空时表示只需对应鉴权身份。",
|
||||
"",
|
||||
(
|
||||
"| # | Method/path | Java Controller.handler | 请求面 | 响应类型 | 鉴权 / 权限 | "
|
||||
"DB/Redis/文件/外部副作用 | 实现与测试状态 |"
|
||||
),
|
||||
"|---:|---|---|---|---|---|---|---|",
|
||||
]
|
||||
)
|
||||
for index, route in enumerate(routes, start=1):
|
||||
declaration = _declaration(route)
|
||||
permission = _permission(route) or "—"
|
||||
auth = {
|
||||
"anonymous": "匿名",
|
||||
"database-token": "DB Token",
|
||||
"server-secret": "server-secret",
|
||||
}[route["auth"]]
|
||||
values = [
|
||||
f"J{index:03d}",
|
||||
f"`{route['method']} {route['path']}`",
|
||||
f"`{route['controller']}.{route['handler']}`",
|
||||
_request_surface(route, declaration),
|
||||
_response_type(route, declaration),
|
||||
f"{auth} / `{permission}`" if permission != "—" else f"{auth} / —",
|
||||
_side_effect(route),
|
||||
_verification(route),
|
||||
]
|
||||
lines.append("| " + " | ".join(_escape(value) for value in values) + " |")
|
||||
|
||||
lines.extend(
|
||||
[
|
||||
"",
|
||||
"## 已观测差异与未覆盖面",
|
||||
"",
|
||||
"- 154 条安全请求面差分最终全部一致。首轮曾发现 5 个空 Body 映射差异;修复 FastAPI 对",
|
||||
" Spring `HttpMessageNotReadableException` 的 code-500 语义后,重新从零执行才得到 154/154。",
|
||||
"- 154 条认证业务面差分最终全部一致;该轮使用有效鉴权与安全业务/校验输入,在不主动成功",
|
||||
" 写入的前提下逐路由对照。证据是 `authenticated-route-results.json`,渲染器会在结果不是",
|
||||
" 154/154、存在失败或跳过时硬失败。",
|
||||
"- 2026-07-20 的隔离差分报告未在 49 个 checks 中观测到响应/所选 headers/数据库副作用",
|
||||
" 不一致;证据是 `main/manager-api-fastapi/compatibility/contract-results.json`,不是人工推断。",
|
||||
"- Hibernate Validator 的 `ConstraintViolation Set` 首条消息无稳定顺序;模型 provider 必填",
|
||||
" 用例比较“消息属于 Java 声明约束集合”与相同错误码,而不伪造一个固定顺序。",
|
||||
"- OTA 时间戳/token 是动态值,差分先比较归一化结构,再分别校验两端 HMAC/Base64 密码学",
|
||||
" 有效性;这属于有意的测试归一化,不是声称字节恒等。",
|
||||
"- 深度差分未直接命中的 133 条中,J125 是下载链路间接覆盖,另 132 条标为 `差分—`;",
|
||||
" 它们有请求面、认证业务面与所属领域测试,但尚无逐路由成功+主要错误+副作用深度对照,不能据此宣称",
|
||||
" 每一种业务状态均已逐接口行为等价。",
|
||||
"- FastAPI 额外提供上述 3 条消费者兼容路由与 live/ready 健康检查;它们没有 Java",
|
||||
" Controller 基线,属于明确、可回退的加法差异。",
|
||||
"- Java 把定时任务放在 Spring 进程;FastAPI 使用独立 jobs 进程和 Redis 分布式锁。这是",
|
||||
" 部署拓扑差异,业务状态和幂等目标保持一致。",
|
||||
"- RAGFlow、阿里云短信、火山语音克隆、真实声纹、真实 LLM、真实 MQTT/MCP/WS 均未用",
|
||||
" 生产凭证联调;自动化只证明 mock 请求格式、超时/错误映射/重试中的已覆盖场景。",
|
||||
"",
|
||||
"## 可复现检查",
|
||||
"",
|
||||
"```bash",
|
||||
"cd main/manager-api-fastapi",
|
||||
".venv/bin/python scripts/extract_java_routes.py --output compatibility/java-routes.json",
|
||||
".venv/bin/python scripts/extract_consumer_routes.py > /tmp/consumer-routes.json",
|
||||
(
|
||||
".venv/bin/pytest -q tests/test_java_route_manifest.py "
|
||||
"tests/test_consumer_route_manifest.py tests/test_compatibility_document.py"
|
||||
),
|
||||
"```",
|
||||
"",
|
||||
"逐接口差分的启动、隔离库、mock 与执行命令见 `docs/manager-api-fastapi-test-report.md`;",
|
||||
"本文件只陈述已落盘的结果,不把缺少真实密钥的外部联调列为通过。",
|
||||
]
|
||||
)
|
||||
rendered = "\n".join(lines) + "\n"
|
||||
if len(re.findall(r"^\| J\d{3} \|", rendered, flags=re.MULTILINE)) != 154:
|
||||
raise AssertionError("rendered route matrix is not closed")
|
||||
return rendered
|
||||
|
||||
|
||||
def main() -> int:
|
||||
parser = argparse.ArgumentParser()
|
||||
parser.add_argument("--output", type=Path)
|
||||
args = parser.parse_args()
|
||||
rendered = render()
|
||||
if args.output is None:
|
||||
print(rendered, end="")
|
||||
else:
|
||||
args.output.parent.mkdir(parents=True, exist_ok=True)
|
||||
args.output.write_text(rendered, encoding="utf-8")
|
||||
return 0
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
raise SystemExit(main())
|
||||
@@ -0,0 +1,220 @@
|
||||
#!/bin/sh
|
||||
set -eu
|
||||
|
||||
SCRIPT_DIR=$(CDPATH= cd -- "$(dirname -- "$0")" && pwd)
|
||||
TARGET_DIR=$(CDPATH= cd -- "${SCRIPT_DIR}/.." && pwd)
|
||||
REPOSITORY_ROOT=$(CDPATH= cd -- "${TARGET_DIR}/../.." && pwd)
|
||||
RUNTIME="${REPOSITORY_ROOT}/.runtime"
|
||||
JAVA_DIR="${REPOSITORY_ROOT}/main/manager-api"
|
||||
STATE_DIR="${TARGET_DIR}/.test-runtime/contract"
|
||||
JAVA_PORT="${CONTRACT_JAVA_PORT:-18082}"
|
||||
FASTAPI_PORT="${CONTRACT_FASTAPI_PORT:-18083}"
|
||||
MOCK_PORT="${CONTRACT_MOCK_PORT:-18084}"
|
||||
JAVA_PID=""
|
||||
FASTAPI_PID=""
|
||||
MOCK_PID=""
|
||||
|
||||
mkdir -p "${STATE_DIR}"
|
||||
|
||||
stop_process() {
|
||||
pid=$1
|
||||
if [ -n "${pid}" ] && kill -0 "${pid}" 2>/dev/null; then
|
||||
children=$(pgrep -P "${pid}" 2>/dev/null || true)
|
||||
if [ -n "${children}" ]; then
|
||||
kill ${children} 2>/dev/null || true
|
||||
fi
|
||||
kill "${pid}" 2>/dev/null || true
|
||||
wait "${pid}" 2>/dev/null || true
|
||||
fi
|
||||
}
|
||||
|
||||
cleanup() {
|
||||
stop_process "${FASTAPI_PID}"
|
||||
stop_process "${JAVA_PID}"
|
||||
stop_process "${MOCK_PID}"
|
||||
}
|
||||
trap cleanup EXIT INT TERM
|
||||
|
||||
wait_for_url() {
|
||||
name=$1
|
||||
url=$2
|
||||
attempts=0
|
||||
until curl --fail --silent --show-error "${url}" >/dev/null 2>&1; do
|
||||
attempts=$((attempts + 1))
|
||||
if [ "${attempts}" -ge 240 ]; then
|
||||
echo "Timed out waiting for ${name}; inspect ${STATE_DIR}." >&2
|
||||
return 1
|
||||
fi
|
||||
sleep 0.25
|
||||
done
|
||||
}
|
||||
|
||||
assert_clean_runtime_logs() {
|
||||
for log in "${STATE_DIR}/fastapi.log" "${STATE_DIR}/external-mock.log"; do
|
||||
if rg -n -i 'UnsupportedFieldAttributeWarning|Traceback|\bERROR\b|\bWARNING\b' "${log}"; then
|
||||
echo "Unexpected warning or error in ${log}." >&2
|
||||
return 1
|
||||
fi
|
||||
done
|
||||
|
||||
# Logback's own bootstrap diagnostics contain tokens such as ERROR_FILE.
|
||||
# Runtime application records start with the configured full ISO date.
|
||||
java_errors=$(rg -n '^[0-9]{4}-[0-9]{2}-[0-9]{2} .*ERROR.* - ' "${STATE_DIR}/java.log" || true)
|
||||
if [ -n "${java_errors}" ]; then
|
||||
java_error_count=$(printf '%s\n' "${java_errors}" | wc -l | tr -d ' ')
|
||||
missing_body_errors=$(printf '%s\n' "${java_errors}" | rg -F 'Required request body is missing:' | wc -l | tr -d ' ')
|
||||
missing_device_errors=$(printf '%s\n' "${java_errors}" | rg -F "Required request header 'Device-Id'" | wc -l | tr -d ' ')
|
||||
object_array_errors=$(printf '%s\n' "${java_errors}" | rg -F 'JSON parse error: Cannot deserialize value of type' | wc -l | tr -d ' ')
|
||||
missing_query_errors=$(printf '%s\n' "${java_errors}" | rg -e "Required request parameter '(ids|modelType|id)' for method parameter type String is not present" | wc -l | tr -d ' ')
|
||||
multipart_errors=$(printf '%s\n' "${java_errors}" | rg -F 'Current request is not a multipart request' | wc -l | tr -d ' ')
|
||||
caller_mac_errors=$(printf '%s\n' "${java_errors}" | rg -F 'Cannot invoke "String.toLowerCase()" because "callerMac" is null' | wc -l | tr -d ' ')
|
||||
null_message_errors=$(printf '%s\n' "${java_errors}" | rg -e ' - null$' | wc -l | tr -d ' ')
|
||||
blank_message_errors=$(printf '%s\n' "${java_errors}" | rg -e ' - $' | wc -l | tr -d ' ')
|
||||
if [ "${java_error_count}" -ne 36 ] \
|
||||
|| [ "${missing_body_errors}" -ne 13 ] \
|
||||
|| [ "${missing_device_errors}" -ne 5 ] \
|
||||
|| [ "${object_array_errors}" -ne 8 ] \
|
||||
|| [ "${missing_query_errors}" -ne 4 ] \
|
||||
|| [ "${multipart_errors}" -ne 3 ] \
|
||||
|| [ "${caller_mac_errors}" -ne 1 ] \
|
||||
|| [ "${null_message_errors}" -ne 1 ] \
|
||||
|| [ "${blank_message_errors}" -ne 1 ]; then
|
||||
printf '%s\n' "${java_errors}" >&2
|
||||
echo "The two 154-route safe surfaces did not produce the exact expected Java baseline error profile." >&2
|
||||
return 1
|
||||
fi
|
||||
unexpected_java_errors=$(
|
||||
printf '%s\n' "${java_errors}" |
|
||||
rg -v -e "Required request header 'Device-Id' for method parameter type String is not present" \
|
||||
-e 'Required request body is missing:' \
|
||||
-e 'JSON parse error: Cannot deserialize value of type .* from Object value \(token .*START_OBJECT.*\)' \
|
||||
-e "Required request parameter '(ids|modelType|id)' for method parameter type String is not present" \
|
||||
-e 'Current request is not a multipart request' \
|
||||
-e 'Cannot invoke "String.toLowerCase\(\)" because "callerMac" is null' \
|
||||
-e ' - null$' \
|
||||
-e ' - $' || true
|
||||
)
|
||||
if [ -n "${unexpected_java_errors}" ]; then
|
||||
printf '%s\n' "${unexpected_java_errors}" >&2
|
||||
echo "Unexpected Java baseline error; only the counted safe surface-validation paths are allowed." >&2
|
||||
return 1
|
||||
fi
|
||||
fi
|
||||
}
|
||||
|
||||
start_java() {
|
||||
(
|
||||
cd "${JAVA_DIR}"
|
||||
JAVA_HOME="${RUNTIME}/jdk" \
|
||||
PATH="${RUNTIME}/jdk/bin:${RUNTIME}/maven/bin:${PATH}" \
|
||||
exec "${RUNTIME}/maven/bin/mvn" \
|
||||
-Dmaven.repo.local="${RUNTIME}/m2" \
|
||||
-DskipTests spring-boot:run \
|
||||
-Dspring-boot.run.arguments="--server.port=${JAVA_PORT} \
|
||||
--spring.datasource.druid.url=jdbc:mysql://127.0.0.1:${TEST_MYSQL_PORT}/manager_java_test?useUnicode=true&characterEncoding=UTF-8&serverTimezone=Asia/Shanghai&nullCatalogMeansCurrent=true&allowMultiQueries=true \
|
||||
--spring.datasource.druid.username=xiaozhi_test \
|
||||
--spring.datasource.druid.password=isolated-test-only \
|
||||
--spring.data.redis.host=127.0.0.1 \
|
||||
--spring.data.redis.port=${TEST_REDIS_PORT} \
|
||||
--spring.data.redis.database=1 \
|
||||
--spring.data.redis.password="
|
||||
) >"${STATE_DIR}/java.log" 2>&1 &
|
||||
JAVA_PID=$!
|
||||
wait_for_url "Java baseline" "http://127.0.0.1:${JAVA_PORT}/xiaozhi/ota/"
|
||||
}
|
||||
|
||||
start_mock() {
|
||||
(
|
||||
cd "${TARGET_DIR}"
|
||||
exec .venv/bin/uvicorn tests.compatibility.external_mock:app \
|
||||
--host 127.0.0.1 --port "${MOCK_PORT}" --log-level warning
|
||||
) >"${STATE_DIR}/external-mock.log" 2>&1 &
|
||||
MOCK_PID=$!
|
||||
wait_for_url "external-service mock" "http://127.0.0.1:${MOCK_PORT}/health"
|
||||
}
|
||||
|
||||
start_fastapi() {
|
||||
(
|
||||
cd "${TARGET_DIR}"
|
||||
APP_ENVIRONMENT=test \
|
||||
APP_DATABASE_URL="${TEST_FASTAPI_DATABASE_URL}" \
|
||||
APP_REDIS_URL="${TEST_FASTAPI_REDIS_URL}" \
|
||||
APP_SERVER_SECRET_OVERRIDE=contract-server-secret \
|
||||
APP_UPLOAD_DIR="${STATE_DIR}/uploads" \
|
||||
APP_LOG_LEVEL=WARNING \
|
||||
exec .venv/bin/uvicorn app.main:app \
|
||||
--host 127.0.0.1 --port "${FASTAPI_PORT}" --log-level warning
|
||||
) >"${STATE_DIR}/fastapi.log" 2>&1 &
|
||||
FASTAPI_PID=$!
|
||||
wait_for_url "FastAPI target" "http://127.0.0.1:${FASTAPI_PORT}/xiaozhi/health/ready"
|
||||
}
|
||||
|
||||
cd "${TARGET_DIR}"
|
||||
|
||||
./scripts/isolated-env.sh reset
|
||||
./scripts/isolated-env.sh migrate
|
||||
eval "$(./scripts/isolated-env.sh env)"
|
||||
|
||||
start_mock
|
||||
|
||||
# The retained Java startup creates the SM2 key pair on an empty schema.
|
||||
start_java
|
||||
.venv/bin/python -m tests.compatibility.seed_contract_data \
|
||||
--mysql-port "${TEST_MYSQL_PORT}" --mock-port "${MOCK_PORT}"
|
||||
|
||||
# Reload Java after the deterministic fixture replaced server params. FLUSHALL
|
||||
# is safe here because this is the dedicated Redis on TEST_REDIS_PORT and the
|
||||
# FastAPI target has not started yet.
|
||||
stop_process "${JAVA_PID}"
|
||||
JAVA_PID=""
|
||||
"${RUNTIME}/redis/bin/redis-cli" -h 127.0.0.1 -p "${TEST_REDIS_PORT}" FLUSHALL >/dev/null
|
||||
start_java
|
||||
start_fastapi
|
||||
|
||||
# Restore fixed dates after Java startup and allow no stale async baseline write
|
||||
# to leak into the first read-only comparison.
|
||||
.venv/bin/python -m tests.compatibility.seed_contract_data \
|
||||
--mysql-port "${TEST_MYSQL_PORT}" --mock-port "${MOCK_PORT}"
|
||||
sleep 1
|
||||
.venv/bin/python -m tests.compatibility.seed_contract_data \
|
||||
--mysql-port "${TEST_MYSQL_PORT}" --mock-port "${MOCK_PORT}"
|
||||
|
||||
TEST_FASTAPI_DATABASE_URL="${TEST_FASTAPI_DATABASE_URL}" \
|
||||
TEST_FASTAPI_REDIS_URL="${TEST_FASTAPI_REDIS_URL}" \
|
||||
APP_DATABASE_URL="${TEST_FASTAPI_DATABASE_URL}" \
|
||||
APP_REDIS_URL="${TEST_FASTAPI_REDIS_URL}" \
|
||||
APP_ENVIRONMENT=test \
|
||||
.venv/bin/pytest -q tests/integration/test_isolated_runtime.py
|
||||
|
||||
.venv/bin/python -m tests.compatibility.route_surface_runner \
|
||||
--java-base "http://127.0.0.1:${JAVA_PORT}/xiaozhi" \
|
||||
--fastapi-base "http://127.0.0.1:${FASTAPI_PORT}/xiaozhi" \
|
||||
--mock-base "http://127.0.0.1:${MOCK_PORT}" \
|
||||
--mysql-port "${TEST_MYSQL_PORT}" \
|
||||
--output compatibility/route-surface-results.json
|
||||
|
||||
.venv/bin/python -m tests.compatibility.authenticated_route_runner \
|
||||
--java-base "http://127.0.0.1:${JAVA_PORT}/xiaozhi" \
|
||||
--fastapi-base "http://127.0.0.1:${FASTAPI_PORT}/xiaozhi" \
|
||||
--mock-base "http://127.0.0.1:${MOCK_PORT}" \
|
||||
--mysql-port "${TEST_MYSQL_PORT}" \
|
||||
--output compatibility/authenticated-route-results.json
|
||||
|
||||
.venv/bin/python -m tests.compatibility.differential_runner \
|
||||
--java-base "http://127.0.0.1:${JAVA_PORT}/xiaozhi" \
|
||||
--fastapi-base "http://127.0.0.1:${FASTAPI_PORT}/xiaozhi" \
|
||||
--mock-base "http://127.0.0.1:${MOCK_PORT}" \
|
||||
--mysql-port "${TEST_MYSQL_PORT}" \
|
||||
--output compatibility/contract-results.json
|
||||
|
||||
.venv/bin/python -m tests.compatibility.seed_contract_data \
|
||||
--mysql-port "${TEST_MYSQL_PORT}" --mock-port "${MOCK_PORT}"
|
||||
.venv/bin/python -m tests.compatibility.performance_runner \
|
||||
--java-base "http://127.0.0.1:${JAVA_PORT}/xiaozhi" \
|
||||
--fastapi-base "http://127.0.0.1:${FASTAPI_PORT}/xiaozhi" \
|
||||
--requests 60 --concurrency 6 --warmup 10 \
|
||||
--output compatibility/performance-results.json
|
||||
|
||||
assert_clean_runtime_logs
|
||||
|
||||
echo "Isolated integration, 154-route unauthenticated and authenticated surfaces, deep differential, and performance tests passed."
|
||||
+51
@@ -0,0 +1,51 @@
|
||||
#!/bin/sh
|
||||
set -eu
|
||||
|
||||
: "${LIQUIBASE_URL:?Set LIQUIBASE_URL to the isolated or deployment JDBC URL}"
|
||||
: "${LIQUIBASE_USERNAME:?Set LIQUIBASE_USERNAME}"
|
||||
: "${LIQUIBASE_PASSWORD:?Set LIQUIBASE_PASSWORD}"
|
||||
|
||||
JDBC_URL=${LIQUIBASE_URL}
|
||||
JDBC_USERNAME=${LIQUIBASE_USERNAME}
|
||||
JDBC_PASSWORD=${LIQUIBASE_PASSWORD}
|
||||
unset LIQUIBASE_URL LIQUIBASE_USERNAME LIQUIBASE_PASSWORD
|
||||
export MIGRATION_JDBC_URL=${JDBC_URL}
|
||||
export MIGRATION_USERNAME=${JDBC_USERNAME}
|
||||
export MIGRATION_PASSWORD=${JDBC_PASSWORD}
|
||||
|
||||
SCRIPT_DIR=$(CDPATH= cd -- "$(dirname -- "$0")" && pwd)
|
||||
PROJECT_DIR=$(CDPATH= cd -- "${SCRIPT_DIR}/.." && pwd)
|
||||
POM="${MIGRATION_POM:-${PROJECT_DIR}/migration-pom.xml}"
|
||||
JAVA_RESOURCES="${JAVA_RESOURCES_DIR:-${PROJECT_DIR}/../manager-api/src/main/resources}"
|
||||
if [ -n "${MIGRATION_RUNNER_JAR:-}" ]; then
|
||||
exec java -jar "${MIGRATION_RUNNER_JAR}"
|
||||
fi
|
||||
|
||||
RUNTIME_ROOT="${PROJECT_DIR}/../../.runtime"
|
||||
if [ -n "${MAVEN_BIN:-}" ]; then
|
||||
MAVEN="${MAVEN_BIN}"
|
||||
elif [ -x "${RUNTIME_ROOT}/maven/bin/mvn" ]; then
|
||||
MAVEN="${RUNTIME_ROOT}/maven/bin/mvn"
|
||||
else
|
||||
MAVEN="mvn"
|
||||
fi
|
||||
|
||||
MAVEN_REPOSITORY_ARGS=""
|
||||
if [ -n "${MAVEN_LOCAL_REPOSITORY:-}" ]; then
|
||||
MAVEN_REPOSITORY_ARGS="-Dmaven.repo.local=${MAVEN_LOCAL_REPOSITORY}"
|
||||
elif [ -x "${RUNTIME_ROOT}/maven/bin/mvn" ]; then
|
||||
MAVEN_REPOSITORY_ARGS="-Dmaven.repo.local=${RUNTIME_ROOT}/m2"
|
||||
fi
|
||||
|
||||
"${MAVEN}" -B -f "${POM}" ${MAVEN_REPOSITORY_ARGS} \
|
||||
-Djava.resources.dir="${JAVA_RESOURCES}" package
|
||||
if [ -n "${JAVA_BIN:-}" ]; then
|
||||
JAVA="${JAVA_BIN}"
|
||||
elif [ -n "${JAVA_HOME:-}" ] && [ -x "${JAVA_HOME}/bin/java" ]; then
|
||||
JAVA="${JAVA_HOME}/bin/java"
|
||||
elif [ -x "${RUNTIME_ROOT}/jdk/bin/java" ]; then
|
||||
JAVA="${RUNTIME_ROOT}/jdk/bin/java"
|
||||
else
|
||||
JAVA="java"
|
||||
fi
|
||||
exec "${JAVA}" -jar "${PROJECT_DIR}/target/manager-api-liquibase-runner-1.0.0-all.jar"
|
||||
+20
@@ -0,0 +1,20 @@
|
||||
#!/bin/sh
|
||||
set -eu
|
||||
|
||||
SCRIPT_DIR=$(CDPATH= cd -- "$(dirname -- "$0")" && pwd)
|
||||
PROJECT_DIR=$(CDPATH= cd -- "${SCRIPT_DIR}/.." && pwd)
|
||||
PYTHON="${PYTHON_BIN:-${PROJECT_DIR}/.venv/bin/python}"
|
||||
|
||||
if [ ! -x "${PYTHON}" ]; then
|
||||
echo "Python environment is missing; run 'uv sync --locked' in ${PROJECT_DIR}" >&2
|
||||
exit 1
|
||||
fi
|
||||
|
||||
cd "${PROJECT_DIR}"
|
||||
exec "${PYTHON}" -m uvicorn app.main:app \
|
||||
--host "${APP_HOST:-0.0.0.0}" \
|
||||
--port "${APP_PORT:-8002}" \
|
||||
--workers "${APP_WORKERS:-1}" \
|
||||
--timeout-graceful-shutdown "${APP_GRACEFUL_SHUTDOWN_SECONDS:-30}" \
|
||||
--proxy-headers \
|
||||
--forwarded-allow-ips "${APP_FORWARDED_ALLOW_IPS:-127.0.0.1}"
|
||||
+14
@@ -0,0 +1,14 @@
|
||||
#!/bin/sh
|
||||
set -eu
|
||||
|
||||
SCRIPT_DIR=$(CDPATH= cd -- "$(dirname -- "$0")" && pwd)
|
||||
PROJECT_DIR=$(CDPATH= cd -- "${SCRIPT_DIR}/.." && pwd)
|
||||
PYTHON="${PYTHON_BIN:-${PROJECT_DIR}/.venv/bin/python}"
|
||||
|
||||
if [ ! -x "${PYTHON}" ]; then
|
||||
echo "Python environment is missing; run 'uv sync --locked' in ${PROJECT_DIR}" >&2
|
||||
exit 1
|
||||
fi
|
||||
|
||||
cd "${PROJECT_DIR}"
|
||||
exec "${PYTHON}" -m app.jobs.worker
|
||||
@@ -0,0 +1 @@
|
||||
"""Manager API FastAPI automated tests."""
|
||||
@@ -0,0 +1 @@
|
||||
"""Executable Java/FastAPI compatibility-test support."""
|
||||
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user