feat: add FastAPI manager API compatibility baseline

This commit is contained in:
Tyke Chen
2026-07-20 17:00:13 +08:00
parent 7c58fa37b2
commit 804ddb51f2
140 changed files with 47169 additions and 1 deletions
@@ -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)