mirror of
https://github.com/xinnan-tech/xiaozhi-esp32-server.git
synced 2026-07-23 23:53:55 +08:00
feat: add FastAPI manager API compatibility baseline
This commit is contained in:
@@ -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)
|
||||
Reference in New Issue
Block a user