mirror of
https://github.com/xinnan-tech/xiaozhi-esp32-server.git
synced 2026-07-24 08:03:53 +08:00
feat: add FastAPI manager API compatibility baseline
This commit is contained in:
@@ -0,0 +1 @@
|
||||
"""Business services."""
|
||||
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,521 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import base64
|
||||
import hashlib
|
||||
import json
|
||||
import logging
|
||||
import math
|
||||
import urllib.parse
|
||||
from copy import deepcopy
|
||||
from typing import Any, cast
|
||||
|
||||
from cryptography.hazmat.primitives.ciphers import Cipher, algorithms, modes
|
||||
from redis.asyncio import Redis
|
||||
|
||||
from app.core.errors import AppError
|
||||
from app.core.redis import JavaRedisCodec, get_redis
|
||||
from app.repositories.config import ConfigRepository
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
class ConfigService:
|
||||
def __init__(self, repository: ConfigRepository, *, redis: Redis | None = None):
|
||||
self.repository = repository
|
||||
self.redis = redis or get_redis()
|
||||
|
||||
async def get_config(self, *, use_cache: bool) -> dict[str, Any]:
|
||||
if use_cache:
|
||||
cached = JavaRedisCodec.decode(await cast(Any, self.redis.get)("server:config"))
|
||||
if isinstance(cached, dict):
|
||||
return cast(dict[str, Any], cached)
|
||||
result = self._build_base_config(await self.repository.list_params())
|
||||
template = await self.repository.get_default_template()
|
||||
if template is None:
|
||||
raise AppError(10183)
|
||||
await self._build_module_config(
|
||||
result=result,
|
||||
assistant_name=None,
|
||||
prompt=None,
|
||||
summary_memory=None,
|
||||
voice=None,
|
||||
reference_audio=None,
|
||||
reference_text=None,
|
||||
language=None,
|
||||
tts_volume=None,
|
||||
tts_rate=None,
|
||||
tts_pitch=None,
|
||||
vad_model_id=self._string(template.get("vad_model_id")),
|
||||
asr_model_id=self._string(template.get("asr_model_id")),
|
||||
llm_model_id=None,
|
||||
vllm_model_id=None,
|
||||
slm_model_id=None,
|
||||
tts_model_id=None,
|
||||
mem_model_id=None,
|
||||
intent_model_id=None,
|
||||
rag_model_id=None,
|
||||
)
|
||||
await cast(Any, self.redis.set)("server:config", JavaRedisCodec.encode(result), ex=24 * 60 * 60)
|
||||
return result
|
||||
|
||||
async def get_agent_models(
|
||||
self,
|
||||
mac_address: str,
|
||||
selected_module: dict[str, str],
|
||||
) -> dict[str, Any]:
|
||||
temporary_key = f"tmp_register_mac:{mac_address}"
|
||||
temporary = JavaRedisCodec.decode(await cast(Any, self.redis.get)(temporary_key))
|
||||
if temporary == "true":
|
||||
await cast(Any, self.redis.delete)(temporary_key)
|
||||
return await self.get_config(use_cache=True)
|
||||
|
||||
device = await self.repository.get_device_by_mac(mac_address)
|
||||
if device is None:
|
||||
safe_address = mac_address.replace(":", "_").lower()
|
||||
activation = JavaRedisCodec.decode(
|
||||
await cast(Any, self.redis.get)(f"ota:activation:data:{safe_address}")
|
||||
)
|
||||
if isinstance(activation, dict) and activation.get("activation_code"):
|
||||
raise AppError(10042, params=(str(activation["activation_code"]),))
|
||||
raise AppError(10041)
|
||||
|
||||
agent_id = self._string(device.get("agent_id"))
|
||||
agent = await self.repository.get_agent(agent_id or "") if agent_id else None
|
||||
if agent is None:
|
||||
raise AppError(10053)
|
||||
|
||||
voice: str | None = None
|
||||
reference_audio: str | None = None
|
||||
reference_text: str | None = None
|
||||
language: str | None = None
|
||||
voice_id = self._string(agent.get("tts_voice_id"))
|
||||
timbre = await self._timbre(voice_id) if voice_id else None
|
||||
if timbre is not None:
|
||||
voice = self._string(timbre.get("tts_voice"))
|
||||
reference_audio = self._string(timbre.get("reference_audio"))
|
||||
reference_text = self._string(timbre.get("reference_text"))
|
||||
chosen_language = self._string(agent.get("tts_language"))
|
||||
if chosen_language and chosen_language.strip():
|
||||
language = chosen_language
|
||||
else:
|
||||
languages = self._string(timbre.get("languages"))
|
||||
if languages and languages.strip():
|
||||
language = languages.split("、", 1)[0].strip()
|
||||
elif voice_id:
|
||||
clone = await self.repository.get_voice_clone(voice_id)
|
||||
if clone is not None:
|
||||
voice = self._string(clone.get("voice_id"))
|
||||
chosen_language = self._string(agent.get("tts_language"))
|
||||
language = chosen_language if chosen_language and chosen_language.strip() else "普通话"
|
||||
|
||||
result: dict[str, Any] = {
|
||||
"device_max_output_size": await self._param("device_max_output_size", from_cache=True)
|
||||
}
|
||||
memory_model = self._string(agent.get("mem_model_id"))
|
||||
chat_history = agent.get("chat_history_conf")
|
||||
if memory_model == "Memory_nomem":
|
||||
chat_history = 0
|
||||
elif memory_model is not None and memory_model != "Memory_nomem" and chat_history is None:
|
||||
chat_history = 2
|
||||
result["chat_history_conf"] = chat_history
|
||||
|
||||
vad_model_id = self._string(agent.get("vad_model_id"))
|
||||
asr_model_id = self._string(agent.get("asr_model_id"))
|
||||
if selected_module.get("VAD") == vad_model_id:
|
||||
vad_model_id = None
|
||||
if selected_module.get("ASR") == asr_model_id:
|
||||
asr_model_id = None
|
||||
|
||||
if self._string(agent.get("intent_model_id")) != "Intent_nointent":
|
||||
plugins = await self._plugins(str(agent["id"]))
|
||||
if plugins:
|
||||
result["plugins"] = plugins
|
||||
|
||||
mcp_endpoint = await self._mcp_address(str(agent["id"]))
|
||||
if mcp_endpoint and mcp_endpoint.startswith("ws"):
|
||||
result["mcp_endpoint"] = mcp_endpoint.replace("/mcp/", "/call/")
|
||||
|
||||
context_providers = self._json_value(await self.repository.get_context_providers(str(agent["id"])))
|
||||
if isinstance(context_providers, list) and context_providers:
|
||||
result["context_providers"] = context_providers
|
||||
|
||||
await self._add_voiceprint(str(agent["id"]), result)
|
||||
await self._build_module_config(
|
||||
result=result,
|
||||
assistant_name=self._string(agent.get("agent_name")),
|
||||
prompt=self._string(agent.get("system_prompt")),
|
||||
summary_memory=self._string(agent.get("summary_memory")),
|
||||
voice=voice,
|
||||
reference_audio=reference_audio,
|
||||
reference_text=reference_text,
|
||||
language=language,
|
||||
tts_volume=self._integer(agent.get("tts_volume")),
|
||||
tts_rate=self._integer(agent.get("tts_rate")),
|
||||
tts_pitch=self._integer(agent.get("tts_pitch")),
|
||||
vad_model_id=vad_model_id,
|
||||
asr_model_id=asr_model_id,
|
||||
llm_model_id=self._string(agent.get("llm_model_id")),
|
||||
vllm_model_id=self._string(agent.get("vllm_model_id")),
|
||||
slm_model_id=self._string(agent.get("slm_model_id")),
|
||||
tts_model_id=self._string(agent.get("tts_model_id")),
|
||||
mem_model_id=memory_model,
|
||||
intent_model_id=self._string(agent.get("intent_model_id")),
|
||||
rag_model_id=None,
|
||||
)
|
||||
return result
|
||||
|
||||
async def get_correct_words(self, mac_address: str) -> list[str]:
|
||||
device = await self.repository.get_device_by_mac(mac_address)
|
||||
if device is None or device.get("agent_id") is None:
|
||||
return []
|
||||
rows = await self.repository.get_correct_word_items(str(device["agent_id"]))
|
||||
return [
|
||||
f"{self._java_string(row.get('source_word'))}|{self._java_string(row.get('target_word'))}"
|
||||
for row in rows
|
||||
]
|
||||
|
||||
@staticmethod
|
||||
def _build_base_config(rows: list[dict[str, Any]]) -> dict[str, Any]:
|
||||
config: dict[str, Any] = {}
|
||||
for row in rows:
|
||||
code = str(row.get("param_code") or "")
|
||||
keys = code.split(".")
|
||||
current = config
|
||||
for key in keys[:-1]:
|
||||
if key not in current:
|
||||
current[key] = {}
|
||||
nested = current[key]
|
||||
if not isinstance(nested, dict):
|
||||
raise TypeError(f"configuration path {code} collides with scalar key {key}")
|
||||
current = nested
|
||||
value = str(row.get("param_value") or "")
|
||||
value_type = str(row.get("value_type") or "string").lower()
|
||||
current[keys[-1]] = ConfigService._typed_param(value, value_type)
|
||||
return config
|
||||
|
||||
@staticmethod
|
||||
def _typed_param(value: str, value_type: str) -> Any:
|
||||
if value_type == "number":
|
||||
try:
|
||||
number = float(value)
|
||||
# Java's implementation returns an Integer only when the double
|
||||
# equals its narrowing conversion to a signed 32-bit int.
|
||||
if math.isnan(number):
|
||||
narrowed = 0
|
||||
elif number >= 2**31 - 1:
|
||||
narrowed = 2**31 - 1
|
||||
elif number <= -(2**31):
|
||||
narrowed = -(2**31)
|
||||
else:
|
||||
narrowed = int(number)
|
||||
return narrowed if number == narrowed else number
|
||||
except ValueError:
|
||||
return value
|
||||
if value_type == "boolean":
|
||||
return value.lower() == "true"
|
||||
if value_type == "array":
|
||||
return [item.strip() for item in value.split(";") if item.strip()]
|
||||
if value_type == "json":
|
||||
try:
|
||||
return json.loads(value)
|
||||
except json.JSONDecodeError:
|
||||
return value
|
||||
return value
|
||||
|
||||
async def _build_module_config(
|
||||
self,
|
||||
*,
|
||||
result: dict[str, Any],
|
||||
assistant_name: str | None,
|
||||
prompt: str | None,
|
||||
summary_memory: str | None,
|
||||
voice: str | None,
|
||||
reference_audio: str | None,
|
||||
reference_text: str | None,
|
||||
language: str | None,
|
||||
tts_volume: int | None,
|
||||
tts_rate: int | None,
|
||||
tts_pitch: int | None,
|
||||
vad_model_id: str | None,
|
||||
asr_model_id: str | None,
|
||||
llm_model_id: str | None,
|
||||
vllm_model_id: str | None,
|
||||
slm_model_id: str | None,
|
||||
tts_model_id: str | None,
|
||||
mem_model_id: str | None,
|
||||
intent_model_id: str | None,
|
||||
rag_model_id: str | None,
|
||||
) -> None:
|
||||
selected: dict[str, str] = {}
|
||||
model_types = ("VAD", "ASR", "TTS", "Memory", "Intent", "LLM", "VLLM", "SLM", "RAG")
|
||||
model_ids = (
|
||||
vad_model_id,
|
||||
asr_model_id,
|
||||
tts_model_id,
|
||||
mem_model_id,
|
||||
intent_model_id,
|
||||
llm_model_id,
|
||||
vllm_model_id,
|
||||
slm_model_id,
|
||||
rag_model_id,
|
||||
)
|
||||
intent_llm_id: str | None = None
|
||||
memory_llm_id: str | None = None
|
||||
for model_type, model_id in zip(model_types, model_ids, strict=True):
|
||||
if model_id is None:
|
||||
continue
|
||||
model = await self._model(model_id)
|
||||
if model is None:
|
||||
continue
|
||||
configuration = self._json_value(model.get("config_json"))
|
||||
type_config: dict[str, Any] = {}
|
||||
if isinstance(configuration, dict):
|
||||
configuration = deepcopy(configuration)
|
||||
type_config[str(model["id"])] = configuration
|
||||
if model_type == "TTS":
|
||||
optional_values = {
|
||||
"private_voice": voice,
|
||||
"ref_audio": reference_audio,
|
||||
"ref_text": reference_text,
|
||||
"language": language,
|
||||
"ttsVolume": tts_volume,
|
||||
"ttsRate": tts_rate,
|
||||
"ttsPitch": tts_pitch,
|
||||
}
|
||||
configuration.update({key: value for key, value in optional_values.items() if value is not None})
|
||||
if configuration.get("type") == "huoshan_double_stream" and voice and voice.startswith("S_"):
|
||||
configuration["resource_id"] = "seed-icl-1.0"
|
||||
elif model_type == "Intent":
|
||||
if configuration.get("type") == "intent_llm":
|
||||
intent_llm_id = self._string(configuration.get("llm"))
|
||||
if intent_llm_id == llm_model_id:
|
||||
intent_llm_id = None
|
||||
functions = configuration.get("functions")
|
||||
if isinstance(functions, str) and functions.strip():
|
||||
configuration["functions"] = functions.split(";")
|
||||
elif model_type == "Memory" and configuration.get("type") == "mem_local_short":
|
||||
memory_llm_id = self._string(configuration.get("llm"))
|
||||
if memory_llm_id == llm_model_id:
|
||||
memory_llm_id = None
|
||||
elif model_type == "LLM":
|
||||
for extra_id in (intent_llm_id, memory_llm_id):
|
||||
if extra_id and extra_id not in type_config:
|
||||
extra = await self._model(extra_id)
|
||||
if extra is not None:
|
||||
type_config[str(extra["id"])] = deepcopy(self._json_value(extra.get("config_json")))
|
||||
if slm_model_id and slm_model_id != llm_model_id and slm_model_id not in type_config:
|
||||
small = await self._model(slm_model_id)
|
||||
small_config = None if small is None else self._json_value(small.get("config_json"))
|
||||
if small is not None and small_config is not None:
|
||||
type_config[str(small["id"])] = deepcopy(small_config)
|
||||
result[model_type] = type_config
|
||||
selected[model_type] = str(model["id"])
|
||||
result["selected_module"] = selected
|
||||
if prompt and prompt.strip():
|
||||
replacement = assistant_name if assistant_name and assistant_name.strip() else "小智"
|
||||
prompt = prompt.replace("{{assistant_name}}", replacement)
|
||||
result["prompt"] = prompt
|
||||
result["summaryMemory"] = summary_memory
|
||||
|
||||
async def _plugins(self, agent_id: str) -> dict[str, Any]:
|
||||
mappings = await self.repository.get_plugin_mappings(agent_id)
|
||||
result: dict[str, Any] = {}
|
||||
knowledge_groups: dict[str, list[dict[str, Any]]] = {}
|
||||
knowledge_models: dict[str, dict[str, Any]] = {}
|
||||
for mapping in mappings:
|
||||
provider_code = self._string(mapping.get("provider_code"))
|
||||
if provider_code and provider_code.strip():
|
||||
value = mapping.get("param_info")
|
||||
result[provider_code] = (
|
||||
json.dumps(value, ensure_ascii=False, separators=(",", ":")) if isinstance(value, dict) else value
|
||||
)
|
||||
# Java removes knowledge mappings by iterating the original list backwards, which reverses dataset order.
|
||||
for mapping in reversed(mappings):
|
||||
provider_code = self._string(mapping.get("provider_code"))
|
||||
if provider_code and provider_code.strip():
|
||||
continue
|
||||
dataset = await self.repository.get_dataset(str(mapping["plugin_id"]))
|
||||
if dataset is None or dataset.get("rag_model_id") is None:
|
||||
continue
|
||||
model = await self._model(str(dataset["rag_model_id"]))
|
||||
if model is None or not model.get("model_code"):
|
||||
continue
|
||||
code = str(model["model_code"])
|
||||
knowledge_groups.setdefault(code, []).append(dataset)
|
||||
knowledge_models[code] = model
|
||||
for code, datasets in knowledge_groups.items():
|
||||
model_config = self._json_value(knowledge_models[code].get("config_json"))
|
||||
if not isinstance(model_config, dict):
|
||||
continue
|
||||
names = ",".join(self._java_string(dataset.get("name")) for dataset in datasets)
|
||||
descriptions = ",".join(
|
||||
self._java_string(dataset.get("description")) for dataset in datasets
|
||||
)
|
||||
params = {
|
||||
"base_url": model_config.get("base_url"),
|
||||
"api_key": model_config.get("api_key"),
|
||||
"dataset_ids": [dataset.get("dataset_id") for dataset in datasets],
|
||||
"description": (
|
||||
f"如果用户询问与【{names}】涵盖的主体范围相关内容时应调用本方法,"
|
||||
f"用于查询:{descriptions}"
|
||||
),
|
||||
}
|
||||
result[f"search_from_{code}"] = json.dumps(params, ensure_ascii=False, separators=(",", ":"))
|
||||
return result
|
||||
|
||||
async def _mcp_address(self, agent_id: str) -> str | None:
|
||||
endpoint = await self._param("server.mcp_endpoint", from_cache=True)
|
||||
if endpoint is None or not endpoint.strip() or endpoint == "null":
|
||||
return None
|
||||
parsed = urllib.parse.urlsplit(endpoint)
|
||||
query = parsed.query
|
||||
marker_index = query.find("key=")
|
||||
key = query[marker_index + len("key=") :]
|
||||
scheme = "wss" if parsed.scheme == "https" else "ws"
|
||||
path = parsed.path
|
||||
prefix_path = path[: path.rfind("/")] if "/" in path else ""
|
||||
prefix = urllib.parse.urlunsplit((scheme, parsed.netloc, prefix_path, "", ""))
|
||||
token = self._aes_encrypt(
|
||||
key,
|
||||
json.dumps(
|
||||
{"agentId": hashlib.md5(agent_id.encode(), usedforsecurity=False).hexdigest()},
|
||||
ensure_ascii=False,
|
||||
separators=(", ", ": "),
|
||||
),
|
||||
)
|
||||
return f"{prefix}/mcp/?token={urllib.parse.quote_plus(token)}"
|
||||
|
||||
@staticmethod
|
||||
def _aes_encrypt(key: str, plaintext: str) -> str:
|
||||
key_bytes = key.encode()
|
||||
if len(key_bytes) not in {16, 24, 32}:
|
||||
key_bytes = (key_bytes + bytes(32))[:32]
|
||||
block_size = 16
|
||||
padding_length = block_size - len(plaintext.encode()) % block_size
|
||||
padded = plaintext.encode() + bytes([padding_length]) * padding_length
|
||||
# Java's published MCP token format is AES/ECB/PKCS5Padding; changing modes breaks existing servers.
|
||||
encryptor = Cipher(algorithms.AES(key_bytes), modes.ECB()).encryptor() # noqa: S305
|
||||
encrypted = encryptor.update(padded) + encryptor.finalize()
|
||||
return base64.b64encode(encrypted).decode("ascii")
|
||||
|
||||
async def _add_voiceprint(self, agent_id: str, result: dict[str, Any]) -> None:
|
||||
try:
|
||||
url = await self._param("server.voice_print", from_cache=True)
|
||||
if url is None or not url.strip() or url == "null":
|
||||
return
|
||||
rows = await self.repository.get_voiceprints(agent_id)
|
||||
if not rows:
|
||||
return
|
||||
speakers = [
|
||||
(
|
||||
f"{self._java_string(row.get('id'))},"
|
||||
f"{self._java_string(row.get('source_name'))},{row.get('introduce') or ''}"
|
||||
)
|
||||
for row in rows
|
||||
]
|
||||
threshold_value = await self._param("server.voiceprint_similarity_threshold", from_cache=True)
|
||||
try:
|
||||
threshold = (
|
||||
float(threshold_value)
|
||||
if threshold_value is not None and threshold_value not in ("", "null")
|
||||
else 0.4
|
||||
)
|
||||
except ValueError:
|
||||
threshold = 0.4
|
||||
result["voiceprint"] = {"url": url, "speakers": speakers, "similarity_threshold": threshold}
|
||||
except Exception:
|
||||
logger.warning("Voiceprint configuration lookup failed", exc_info=True)
|
||||
|
||||
async def _param(self, code: str, *, from_cache: bool) -> str | None:
|
||||
if from_cache:
|
||||
cached = JavaRedisCodec.decode(await cast(Any, self.redis.hget)("sys:params", code))
|
||||
if cached is not None:
|
||||
return str(cached)
|
||||
value = await self.repository.get_param_value(code)
|
||||
if value is not None and from_cache:
|
||||
await cast(Any, self.redis.hset)("sys:params", code, JavaRedisCodec.encode(value))
|
||||
await cast(Any, self.redis.expire)("sys:params", 24 * 60 * 60)
|
||||
return value
|
||||
|
||||
async def _model(self, model_id: str) -> dict[str, Any] | None:
|
||||
key = f"model:data:{model_id}"
|
||||
cached = JavaRedisCodec.decode(await cast(Any, self.redis.get)(key))
|
||||
if isinstance(cached, dict):
|
||||
return self._normalize_cached(cast(dict[str, Any], cached))
|
||||
model = await self.repository.get_model(model_id)
|
||||
if model is not None:
|
||||
raw_configuration = model.get("config_json")
|
||||
if isinstance(raw_configuration, str):
|
||||
parsed_configuration = json.loads(raw_configuration)
|
||||
if parsed_configuration is not None and not isinstance(parsed_configuration, dict):
|
||||
raise TypeError("ModelConfigEntity.configJson must be a JSON object")
|
||||
model["config_json"] = parsed_configuration
|
||||
await cast(Any, self.redis.set)(
|
||||
key,
|
||||
JavaRedisCodec.encode(
|
||||
model,
|
||||
java_type="xiaozhi.modules.model.entity.ModelConfigEntity",
|
||||
field_java_types={
|
||||
"configJson": "cn.hutool.json.JSONObject",
|
||||
"creator": "java.lang.Long",
|
||||
"updater": "java.lang.Long",
|
||||
},
|
||||
),
|
||||
ex=24 * 60 * 60,
|
||||
)
|
||||
return model
|
||||
|
||||
async def _timbre(self, timbre_id: str) -> dict[str, Any] | None:
|
||||
key = f"timbre:details:{timbre_id}"
|
||||
cached = JavaRedisCodec.decode(await cast(Any, self.redis.get)(key))
|
||||
if isinstance(cached, dict):
|
||||
return self._normalize_cached(cast(dict[str, Any], cached))
|
||||
timbre = await self.repository.get_timbre(timbre_id)
|
||||
if timbre is not None:
|
||||
await cast(Any, self.redis.set)(
|
||||
key,
|
||||
JavaRedisCodec.encode(
|
||||
timbre,
|
||||
java_type="xiaozhi.modules.timbre.vo.TimbreDetailsVO",
|
||||
field_java_types={"sort": "java.lang.Long"},
|
||||
),
|
||||
ex=24 * 60 * 60,
|
||||
)
|
||||
return timbre
|
||||
|
||||
@staticmethod
|
||||
def _normalize_cached(value: dict[str, Any]) -> dict[str, Any]:
|
||||
aliases = {
|
||||
"modelType": "model_type",
|
||||
"modelCode": "model_code",
|
||||
"modelName": "model_name",
|
||||
"configJson": "config_json",
|
||||
"ttsVoice": "tts_voice",
|
||||
"referenceAudio": "reference_audio",
|
||||
"referenceText": "reference_text",
|
||||
"ttsModelId": "tts_model_id",
|
||||
}
|
||||
return {aliases.get(key, key): item for key, item in value.items() if key != "@class"}
|
||||
|
||||
@staticmethod
|
||||
def _json_value(value: Any) -> Any:
|
||||
if isinstance(value, bytes):
|
||||
value = value.decode()
|
||||
if isinstance(value, str):
|
||||
try:
|
||||
return json.loads(value)
|
||||
except json.JSONDecodeError:
|
||||
return value
|
||||
return value
|
||||
|
||||
@staticmethod
|
||||
def _string(value: Any) -> str | None:
|
||||
return None if value is None else str(value)
|
||||
|
||||
@staticmethod
|
||||
def _integer(value: Any) -> int | None:
|
||||
return None if value is None else int(value)
|
||||
|
||||
@staticmethod
|
||||
def _java_string(value: Any) -> str:
|
||||
return "null" if value is None else str(value)
|
||||
@@ -0,0 +1,133 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import uuid
|
||||
from typing import Any
|
||||
|
||||
from app.core.errors import AppError
|
||||
from app.core.security import AuthUser, shanghai_now_naive
|
||||
from app.repositories.correctword import CorrectWordRepository
|
||||
from app.schemas.correctword import CorrectWordFileBody
|
||||
|
||||
|
||||
def _parse_lines(lines: list[str]) -> list[tuple[str, str]]:
|
||||
result: list[tuple[str, str]] = []
|
||||
for raw in lines:
|
||||
line = raw.strip()
|
||||
if not line or "|" not in line:
|
||||
continue
|
||||
source, target = line.split("|", 1)
|
||||
if source.strip() and target.strip():
|
||||
result.append((source.strip(), target.strip()))
|
||||
return result
|
||||
|
||||
|
||||
def _content_lines(value: str | None) -> list[str]:
|
||||
if value is None:
|
||||
return []
|
||||
# Java String.split keeps one empty element for the empty source string,
|
||||
# while still discarding trailing empty elements for non-empty strings.
|
||||
if value == "":
|
||||
return [""]
|
||||
lines = value.split("\n")
|
||||
while lines and lines[-1] == "":
|
||||
lines.pop()
|
||||
return lines
|
||||
|
||||
|
||||
def file_vo(row: dict[str, Any]) -> dict[str, Any]:
|
||||
return {
|
||||
"id": row.get("id"),
|
||||
"fileName": row.get("file_name"),
|
||||
"wordCount": row.get("word_count"),
|
||||
"content": _content_lines(row.get("content")),
|
||||
"createdAt": row.get("created_at"),
|
||||
"updatedAt": row.get("updated_at"),
|
||||
}
|
||||
|
||||
|
||||
class CorrectWordService:
|
||||
def __init__(self, repository: CorrectWordRepository):
|
||||
self.repository = repository
|
||||
|
||||
@staticmethod
|
||||
def validate(body: CorrectWordFileBody, *, check_size: bool) -> None:
|
||||
if body.file_name is None or not body.file_name.strip():
|
||||
raise AppError(10034, "文件名不能为空")
|
||||
if not body.content:
|
||||
raise AppError(10034, "替换词内容不能为空")
|
||||
if check_size and body.file_size is not None and body.file_size > 1024 * 1024:
|
||||
raise AppError(10204)
|
||||
|
||||
async def create(self, body: CorrectWordFileBody, user: AuthUser) -> dict[str, Any]:
|
||||
self.validate(body, check_size=True)
|
||||
assert body.file_name is not None
|
||||
assert body.content is not None
|
||||
items = _parse_lines(body.content)
|
||||
file_id, now = uuid.uuid4().hex, shanghai_now_naive()
|
||||
values = {
|
||||
"id": file_id,
|
||||
"file_name": body.file_name,
|
||||
"word_count": len(items),
|
||||
"content": "\n".join(body.content),
|
||||
"creator": user.id,
|
||||
"now": now,
|
||||
}
|
||||
async with self.repository.session.begin():
|
||||
if await self.repository.name_exists(user.id, body.file_name):
|
||||
raise AppError(10203)
|
||||
await self.repository.insert_file(values)
|
||||
await self.repository.insert_items(
|
||||
[
|
||||
{"id": uuid.uuid4().hex, "file_id": file_id, "source_word": source, "target_word": target}
|
||||
for source, target in items
|
||||
]
|
||||
)
|
||||
return file_vo({**values, "created_at": now, "updated_at": None})
|
||||
|
||||
async def update(self, file_id: str, body: CorrectWordFileBody, user: AuthUser) -> None:
|
||||
self.validate(body, check_size=False)
|
||||
assert body.file_name is not None
|
||||
assert body.content is not None
|
||||
items = _parse_lines(body.content)
|
||||
async with self.repository.session.begin():
|
||||
row = await self.repository.get_file(file_id, for_update=True)
|
||||
if row is None:
|
||||
return
|
||||
if await self.repository.name_exists(user.id, body.file_name, file_id):
|
||||
raise AppError(500, f"文件名已存在:{body.file_name}")
|
||||
await self.repository.delete_items(file_id)
|
||||
await self.repository.insert_items(
|
||||
[
|
||||
{"id": uuid.uuid4().hex, "file_id": file_id, "source_word": source, "target_word": target}
|
||||
for source, target in items
|
||||
]
|
||||
)
|
||||
await self.repository.update_file(
|
||||
{
|
||||
"id": file_id,
|
||||
"file_name": body.file_name,
|
||||
"word_count": len(items),
|
||||
"content": "\n".join(body.content),
|
||||
"updater": user.id,
|
||||
"now": shanghai_now_naive(),
|
||||
}
|
||||
)
|
||||
|
||||
async def page(self, user: AuthUser, page: str | None, limit: str | None) -> dict[str, Any]:
|
||||
current, size = max(int(page or "1"), 1), int(limit or "10")
|
||||
rows, total = await self.repository.list_files(user.id, offset=(current - 1) * size, limit=size)
|
||||
return {"total": total, "list": [file_vo(row) for row in rows]}
|
||||
|
||||
async def all(self, user: AuthUser) -> list[dict[str, Any]]:
|
||||
rows, _ = await self.repository.list_files(user.id)
|
||||
return [file_vo(row) for row in rows]
|
||||
|
||||
async def get(self, file_id: str) -> dict[str, Any] | None:
|
||||
row = await self.repository.get_file(file_id)
|
||||
return file_vo(row) if row else None
|
||||
|
||||
async def delete(self, file_ids: list[str]) -> None:
|
||||
async with self.repository.session.begin():
|
||||
for file_id in file_ids:
|
||||
if file_id and file_id.strip():
|
||||
await self.repository.delete_file_graph(file_id.strip())
|
||||
@@ -0,0 +1,978 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import base64
|
||||
import hashlib
|
||||
import hmac
|
||||
import json
|
||||
import logging
|
||||
import random
|
||||
import re
|
||||
import secrets
|
||||
import uuid
|
||||
from collections.abc import Callable, Mapping, Sequence
|
||||
from datetime import datetime, timezone
|
||||
from pathlib import Path
|
||||
from typing import Any, cast
|
||||
from zoneinfo import ZoneInfo
|
||||
|
||||
import httpx
|
||||
from redis.asyncio import Redis
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
|
||||
from app.core.config import get_settings
|
||||
from app.core.database import get_session_factory
|
||||
from app.core.errors import AppError
|
||||
from app.core.redis import JavaRedisCodec, get_redis
|
||||
from app.core.security import AuthUser, shanghai_now_naive
|
||||
from app.integrations.mqtt_gateway import post_json
|
||||
from app.repositories.device import DeviceRepository
|
||||
from app.schemas.device import DeviceManualAddRequest, DeviceReportRequest, DeviceUpdateRequest, OtaRecord
|
||||
from app.services.system_params import SystemParamService
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
DEFAULT_TTL_SECONDS = 24 * 60 * 60
|
||||
INVALID_FIRMWARE_URL = (
|
||||
"http://xiaozhi.server.com:8002/xiaozhi/otaMag/download/NOT_ACTIVATED_FIRMWARE_THIS_IS_A_INVALID_URL"
|
||||
)
|
||||
MAC_PATTERN = re.compile(r"^([0-9A-Za-z]{2}[:-]){5}([0-9A-Za-z]{2})$")
|
||||
OTA_ORDER_COLUMNS = {
|
||||
"id": "id",
|
||||
"firmwareName": "firmware_name",
|
||||
"firmware_name": "firmware_name",
|
||||
"type": "type",
|
||||
"version": "version",
|
||||
"size": "size",
|
||||
"sort": "sort",
|
||||
"updateDate": "update_date",
|
||||
"update_date": "update_date",
|
||||
"createDate": "create_date",
|
||||
"create_date": "create_date",
|
||||
}
|
||||
|
||||
|
||||
def is_blank(value: str | None) -> bool:
|
||||
return value is None or not value.strip()
|
||||
|
||||
|
||||
def _java_semicolon_split(value: str) -> list[str]:
|
||||
parts = value.split(";")
|
||||
while parts and parts[-1] == "":
|
||||
parts.pop()
|
||||
return parts
|
||||
|
||||
|
||||
def _mapping(value: Any) -> dict[str, Any] | None:
|
||||
if isinstance(value, dict):
|
||||
if "@class" in value:
|
||||
return {str(key): item for key, item in value.items() if key != "@class"}
|
||||
return {str(key): item for key, item in value.items()}
|
||||
if isinstance(value, list) and len(value) == 2 and isinstance(value[1], dict):
|
||||
return {str(key): item for key, item in value[1].items()}
|
||||
return None
|
||||
|
||||
|
||||
async def redis_get(key: str, client: Redis | None = None) -> Any:
|
||||
selected = client or get_redis()
|
||||
raw = await cast(Any, selected.get(key))
|
||||
return JavaRedisCodec.decode(raw)
|
||||
|
||||
|
||||
async def redis_set(key: str, value: Any, *, ttl: int = DEFAULT_TTL_SECONDS, client: Redis | None = None) -> None:
|
||||
selected = client or get_redis()
|
||||
await cast(Any, selected.set(key, JavaRedisCodec.encode(value), ex=ttl))
|
||||
|
||||
|
||||
async def redis_delete(*keys: str, client: Redis | None = None) -> None:
|
||||
if not keys:
|
||||
return
|
||||
selected = client or get_redis()
|
||||
await cast(Any, selected.delete(*keys))
|
||||
|
||||
|
||||
async def redis_increment(key: str, *, ttl: int = DEFAULT_TTL_SECONDS, client: Redis | None = None) -> int:
|
||||
selected = client or get_redis()
|
||||
value = int(await cast(Any, selected.incr(key)))
|
||||
await cast(Any, selected.expire(key, ttl))
|
||||
return value
|
||||
|
||||
|
||||
class DeviceService:
|
||||
def __init__(
|
||||
self,
|
||||
session: AsyncSession,
|
||||
*,
|
||||
redis_client: Redis | None = None,
|
||||
http_client: httpx.AsyncClient | None = None,
|
||||
):
|
||||
self.session = session
|
||||
self.repository = DeviceRepository(session)
|
||||
self.params = SystemParamService(session)
|
||||
self.redis = redis_client
|
||||
self.http_client = http_client
|
||||
|
||||
async def register_device(self, mac_address: str) -> str:
|
||||
while True:
|
||||
code = f"{secrets.randbelow(1_000_000):06d}"
|
||||
key = f"sys:device:captcha:{code}"
|
||||
if is_blank(cast(str | None, await redis_get(key, self.redis))):
|
||||
await redis_set(key, mac_address, client=self.redis)
|
||||
return code
|
||||
|
||||
async def activate_bound_device(self, *, agent_id: str, activation_code: str, user: AuthUser) -> None:
|
||||
if is_blank(activation_code):
|
||||
raise AppError(10061)
|
||||
code_key = f"ota:activation:code:{activation_code}"
|
||||
device_id_value = await redis_get(code_key, self.redis)
|
||||
if device_id_value in (None, ""):
|
||||
raise AppError(10062)
|
||||
device_id = str(device_id_value)
|
||||
safe_device_id = device_id.replace(":", "_").lower()
|
||||
data_key = f"ota:activation:data:{safe_device_id}"
|
||||
cached = _mapping(await redis_get(data_key, self.redis))
|
||||
if cached is None or str(cached.get("activation_code") or "") != activation_code:
|
||||
raise AppError(10062)
|
||||
if await self.repository.get_device(device_id) is not None:
|
||||
raise AppError(10063)
|
||||
|
||||
now = shanghai_now_naive()
|
||||
values = {
|
||||
"id": device_id,
|
||||
"user_id": user.id,
|
||||
"mac_address": cached.get("mac_address"),
|
||||
"last_connected_at": now,
|
||||
"auto_update": 1,
|
||||
"board": cached.get("board"),
|
||||
"alias": None,
|
||||
"agent_id": agent_id,
|
||||
"app_version": cached.get("app_version"),
|
||||
"sort": None,
|
||||
"updater": user.id,
|
||||
"update_date": now,
|
||||
"creator": user.id,
|
||||
"create_date": now,
|
||||
}
|
||||
try:
|
||||
await self.repository.insert_device(values)
|
||||
await self.session.commit()
|
||||
except Exception:
|
||||
await self.session.rollback()
|
||||
raise
|
||||
await redis_delete(data_key, code_key, f"agent:device:count:{agent_id}", client=self.redis)
|
||||
|
||||
async def list_user_devices(self, user_id: int, agent_id: str) -> list[dict[str, Any]]:
|
||||
devices = await self.repository.get_user_devices(user_id, agent_id)
|
||||
return [self._user_device_view(row) for row in devices]
|
||||
|
||||
async def get_online_data(self, agent_id: str, user: AuthUser) -> str:
|
||||
gateway = await self.params.get_value("server.mqtt_manager_api", from_cache=True)
|
||||
if is_blank(gateway) or gateway == "null":
|
||||
return ""
|
||||
devices = await self.repository.get_user_devices(user.id, agent_id)
|
||||
client_ids = {
|
||||
self._mqtt_client_id(
|
||||
str(device.get("board") or "GID_default"),
|
||||
str(device.get("mac_address") or "unknown"),
|
||||
)
|
||||
for device in devices
|
||||
}
|
||||
if not client_ids:
|
||||
return ""
|
||||
signature_key = await self.params.get_value("server.mqtt_signature_key", from_cache=False)
|
||||
return await post_json(
|
||||
f"http://{gateway}/api/devices/status",
|
||||
{"clientIds": sorted(client_ids)},
|
||||
signature_key or "",
|
||||
timeout_seconds=get_settings().external_request_timeout_seconds,
|
||||
client=self.http_client,
|
||||
)
|
||||
|
||||
async def unbind(self, *, user_id: int, device_id: str) -> None:
|
||||
device = await self.repository.get_device(device_id)
|
||||
if device is None:
|
||||
return
|
||||
mac_address = device.get("mac_address")
|
||||
agent_id = device.get("agent_id")
|
||||
if not is_blank(None if agent_id is None else str(agent_id)):
|
||||
await redis_delete(f"agent:device:count:{agent_id}", client=self.redis)
|
||||
try:
|
||||
await self.repository.delete_device_for_user(device_id, user_id)
|
||||
await self.session.commit()
|
||||
except Exception:
|
||||
await self.session.rollback()
|
||||
raise
|
||||
try:
|
||||
if mac_address is not None:
|
||||
await self.repository.delete_address_books_for_macs([str(mac_address)])
|
||||
await self.session.commit()
|
||||
except Exception:
|
||||
await self.session.rollback()
|
||||
raise
|
||||
await self.refresh_address_book_cache()
|
||||
|
||||
async def update_device(
|
||||
self,
|
||||
*,
|
||||
device_id: str,
|
||||
request: DeviceUpdateRequest,
|
||||
user: AuthUser,
|
||||
) -> bool:
|
||||
device = await self.repository.get_device(device_id)
|
||||
if device is None or int(device.get("user_id") or -1) != user.id:
|
||||
return False
|
||||
await self.repository.update_device_info(
|
||||
device_id,
|
||||
auto_update=request.auto_update,
|
||||
alias=request.alias,
|
||||
updater=user.id,
|
||||
now=shanghai_now_naive(),
|
||||
)
|
||||
await self.session.commit()
|
||||
return True
|
||||
|
||||
async def manual_add(self, *, request: DeviceManualAddRequest, user: AuthUser) -> None:
|
||||
mac_address = request.mac_address
|
||||
if mac_address is not None and await self.repository.get_device_by_mac(mac_address) is not None:
|
||||
raise AppError(10161)
|
||||
now = shanghai_now_naive()
|
||||
values = {
|
||||
"id": uuid.uuid4().hex if mac_address in (None, "") else mac_address,
|
||||
"user_id": user.id,
|
||||
"mac_address": mac_address,
|
||||
"last_connected_at": now,
|
||||
"auto_update": 1,
|
||||
"board": request.board,
|
||||
"alias": None,
|
||||
"agent_id": request.agent_id,
|
||||
"app_version": request.app_version,
|
||||
"sort": None,
|
||||
"updater": user.id,
|
||||
"update_date": now,
|
||||
"creator": user.id,
|
||||
"create_date": now,
|
||||
}
|
||||
try:
|
||||
await self.repository.insert_device(values)
|
||||
await self.session.commit()
|
||||
except Exception:
|
||||
await self.session.rollback()
|
||||
raise
|
||||
agent_cache_id = "null" if request.agent_id is None else request.agent_id
|
||||
await redis_delete(f"agent:device:count:{agent_cache_id}", client=self.redis)
|
||||
|
||||
async def get_tools(self, *, device_id: str, user: AuthUser) -> dict[str, Any] | None:
|
||||
gateway_and_device = await self._gateway_device(device_id, user)
|
||||
if gateway_and_device is None:
|
||||
return None
|
||||
gateway, device = gateway_and_device
|
||||
client_id = self._mqtt_client_id(
|
||||
str(device.get("board") or "GID_default"),
|
||||
str(device.get("mac_address") or "unknown"),
|
||||
)
|
||||
url = f"http://{gateway}/api/commands/{client_id}"
|
||||
all_tools: list[Any] = []
|
||||
cursor: str | None = None
|
||||
while True:
|
||||
params: dict[str, Any] = {"withUserTools": True}
|
||||
if cursor is not None and cursor.strip():
|
||||
params["cursor"] = cursor
|
||||
body = {
|
||||
"type": "mcp",
|
||||
"payload": {"jsonrpc": "2.0", "id": 2, "method": "tools/list", "params": params},
|
||||
}
|
||||
response_body = await self._post_gateway(url, body)
|
||||
if is_blank(response_body):
|
||||
break
|
||||
payload = json.loads(response_body)
|
||||
if not isinstance(payload, dict) or not bool(payload.get("success", False)):
|
||||
break
|
||||
data = payload.get("data")
|
||||
if not isinstance(data, dict):
|
||||
break
|
||||
tools = data.get("tools")
|
||||
if isinstance(tools, list):
|
||||
all_tools.extend(tools)
|
||||
next_cursor = data.get("nextCursor")
|
||||
if not isinstance(next_cursor, str) or not next_cursor.strip():
|
||||
break
|
||||
cursor = next_cursor
|
||||
return None if not all_tools else {"tools": all_tools}
|
||||
|
||||
async def call_tool(
|
||||
self,
|
||||
*,
|
||||
device_id: str,
|
||||
tool_name: str,
|
||||
arguments: dict[str, Any] | None,
|
||||
user: AuthUser,
|
||||
) -> Any:
|
||||
gateway_and_device = await self._gateway_device(device_id, user)
|
||||
if gateway_and_device is None:
|
||||
return None
|
||||
gateway, device = gateway_and_device
|
||||
client_id = self._mqtt_client_id(
|
||||
str(device.get("board") or "GID_default"),
|
||||
str(device.get("mac_address") or "unknown"),
|
||||
)
|
||||
response_body = await self._post_gateway(
|
||||
f"http://{gateway}/api/commands/{client_id}",
|
||||
{
|
||||
"type": "mcp",
|
||||
"payload": {
|
||||
"jsonrpc": "2.0",
|
||||
"id": 2,
|
||||
"method": "tools/call",
|
||||
"params": {"name": tool_name, "arguments": arguments},
|
||||
},
|
||||
},
|
||||
)
|
||||
if is_blank(response_body):
|
||||
return None
|
||||
payload = json.loads(response_body)
|
||||
if not isinstance(payload, dict) or not bool(payload.get("success", False)):
|
||||
return None
|
||||
data = payload.get("data")
|
||||
content = data.get("content") if isinstance(data, dict) else None
|
||||
if not isinstance(content, list) or not content or not isinstance(content[0], dict):
|
||||
return None
|
||||
first = content[0]
|
||||
if first.get("type") != "text" or not isinstance(first.get("text"), str):
|
||||
return None
|
||||
text = str(first["text"])
|
||||
if not text.strip():
|
||||
return None
|
||||
trimmed = text.strip()
|
||||
if trimmed.startswith("{") or trimmed.startswith("["):
|
||||
try:
|
||||
parsed = json.loads(trimmed)
|
||||
return parsed if isinstance(parsed, dict) else trimmed
|
||||
except json.JSONDecodeError:
|
||||
return trimmed
|
||||
if trimmed == "true":
|
||||
return True
|
||||
if trimmed == "false":
|
||||
return False
|
||||
return trimmed
|
||||
|
||||
async def address_book(self, mac_address: str) -> list[dict[str, Any]]:
|
||||
rows = await self.repository.get_address_book(mac_address)
|
||||
for row in rows:
|
||||
if row.get("has_permission") is not None:
|
||||
row["has_permission"] = bool(row["has_permission"])
|
||||
return rows
|
||||
|
||||
async def lookup_address_book(self, *, caller_mac: str, nickname: str) -> dict[str, str | None] | None:
|
||||
books = await self.all_address_books()
|
||||
caller_book = books.get(caller_mac.lower())
|
||||
if caller_book is None:
|
||||
return None
|
||||
target_with_permission = caller_book.get(nickname)
|
||||
if target_with_permission is None:
|
||||
return None
|
||||
parts = target_with_permission.split("|")
|
||||
target_mac = parts[0]
|
||||
has_permission = len(parts) > 1 and parts[1] == "1"
|
||||
target_book = books.get(target_mac.lower())
|
||||
if target_book is None:
|
||||
return None
|
||||
caller_nickname = target_book.get(caller_mac.lower())
|
||||
return {
|
||||
"targetMac": target_mac,
|
||||
"callerNickname": caller_nickname,
|
||||
"hasPermission": "true" if has_permission else "false",
|
||||
}
|
||||
|
||||
async def call_by_nickname(self, *, caller_mac: str, nickname: str, answer: bool) -> dict[str, Any]:
|
||||
books = await self.all_address_books()
|
||||
if answer:
|
||||
return await self._post_call("/api/call/accept", {"mac": caller_mac}, "接听")
|
||||
caller_book = books.get(caller_mac.lower())
|
||||
if caller_book is None or nickname not in caller_book:
|
||||
return {"status": "error", "message": f"未找到备注为'{nickname}'的设备"}
|
||||
parts = caller_book[nickname].split("|")
|
||||
target_mac = parts[0]
|
||||
allowed = len(parts) > 1 and parts[1] == "1"
|
||||
if not allowed:
|
||||
return {"status": "error", "message": "呼叫失败,您没有权限呼叫该设备"}
|
||||
target_book = books.get(target_mac.lower())
|
||||
caller_nickname = target_book.get(caller_mac.lower()) if target_book is not None else None
|
||||
if is_blank(caller_nickname):
|
||||
caller = await self.repository.get_device_by_mac(caller_mac)
|
||||
if caller is None:
|
||||
raise RuntimeError("caller device does not exist")
|
||||
caller_nickname = None if caller.get("alias") is None else str(caller["alias"])
|
||||
if is_blank(caller_nickname):
|
||||
caller_nickname = self._mac_device_name(caller_mac)
|
||||
return await self._post_call(
|
||||
"/api/call/request",
|
||||
{"caller_mac": caller_mac, "target_mac": target_mac, "caller_nickname": caller_nickname},
|
||||
"呼叫",
|
||||
)
|
||||
|
||||
async def save_address_book(
|
||||
self,
|
||||
*,
|
||||
mac_address: str,
|
||||
target_mac: str,
|
||||
alias: str | None,
|
||||
has_permission: bool | None,
|
||||
actor: int,
|
||||
) -> None:
|
||||
record = await self.repository.get_address_book_record(mac_address, target_mac)
|
||||
now = shanghai_now_naive()
|
||||
if record is None:
|
||||
final_alias = alias
|
||||
if is_blank(final_alias):
|
||||
target = await self.repository.get_device_by_mac(target_mac)
|
||||
if target is None:
|
||||
raise RuntimeError("target device does not exist")
|
||||
final_alias = None if target.get("alias") is None else str(target["alias"])
|
||||
final_alias = await self._unique_alias(mac_address, final_alias)
|
||||
await self.repository.insert_address_book(
|
||||
mac_address=mac_address,
|
||||
target_mac=target_mac,
|
||||
alias=final_alias,
|
||||
has_permission=has_permission,
|
||||
actor=actor,
|
||||
now=now,
|
||||
)
|
||||
await self.session.commit()
|
||||
else:
|
||||
if alias is not None:
|
||||
await self.repository.update_address_alias(
|
||||
mac_address,
|
||||
target_mac,
|
||||
await self._unique_alias(mac_address, alias),
|
||||
now=now,
|
||||
)
|
||||
await self.session.commit()
|
||||
await self.refresh_address_book_cache()
|
||||
if has_permission is not None:
|
||||
await self.repository.update_address_permission(
|
||||
mac_address,
|
||||
target_mac,
|
||||
has_permission,
|
||||
now=now,
|
||||
)
|
||||
await self.session.commit()
|
||||
await self.refresh_address_book_cache()
|
||||
|
||||
async def all_address_books(self) -> dict[str, dict[str, str]]:
|
||||
cached = _mapping(await redis_get("device:address_book:all", self.redis))
|
||||
if cached is not None:
|
||||
result: dict[str, dict[str, str]] = {}
|
||||
for key, value in cached.items():
|
||||
nested = _mapping(value)
|
||||
if nested is not None:
|
||||
result[key] = {str(field): str(item) for field, item in nested.items()}
|
||||
return result
|
||||
return await self.refresh_address_book_cache()
|
||||
|
||||
async def refresh_address_book_cache(self) -> dict[str, dict[str, str]]:
|
||||
records = await self.repository.get_all_address_book()
|
||||
result: dict[str, dict[str, str]] = {}
|
||||
reverse: dict[str, str] = {}
|
||||
for record in records:
|
||||
mac_a = str(record["mac_address"]).lower()
|
||||
mac_b = str(record["target_mac"]).lower()
|
||||
alias = record.get("alias")
|
||||
if alias not in (None, ""):
|
||||
alias_string = str(alias)
|
||||
result.setdefault(mac_a, {})[alias_string] = (
|
||||
f"{mac_b}|{'1' if bool(record.get('has_permission')) else '0'}"
|
||||
)
|
||||
reverse[f"{mac_b}:{mac_a}"] = alias_string
|
||||
for record in records:
|
||||
mac_a = str(record["mac_address"]).lower()
|
||||
mac_b = str(record["target_mac"]).lower()
|
||||
reverse_alias = reverse.get(f"{mac_a}:{mac_b}")
|
||||
if isinstance(reverse_alias, str) and reverse_alias:
|
||||
result.setdefault(mac_b, {})[mac_a] = reverse_alias
|
||||
await redis_set("device:address_book:all", result, client=self.redis)
|
||||
return result
|
||||
|
||||
async def check_ota(
|
||||
self,
|
||||
*,
|
||||
device_id: str,
|
||||
client_id: str,
|
||||
report: DeviceReportRequest,
|
||||
request_url: str,
|
||||
client_ip: str,
|
||||
defer_connection_update: Callable[[str, str | None, str | None], None] | None = None,
|
||||
) -> dict[str, Any]:
|
||||
now = datetime.now(tz=ZoneInfo(get_settings().timezone))
|
||||
utc_offset = now.utcoffset()
|
||||
response: dict[str, Any] = {
|
||||
"server_time": {
|
||||
"timestamp": int(now.timestamp() * 1000),
|
||||
"timeZone": get_settings().timezone,
|
||||
"timezone_offset": int((utc_offset.total_seconds() if utc_offset is not None else 0) / 60),
|
||||
},
|
||||
"activation": None,
|
||||
"error": None,
|
||||
"firmware": None,
|
||||
"websocket": None,
|
||||
"mqtt": None,
|
||||
}
|
||||
device = await self.repository.get_device_by_mac(device_id)
|
||||
if device is None:
|
||||
if report.application is None:
|
||||
raise RuntimeError("application is required")
|
||||
response["firmware"] = {
|
||||
"version": report.application.version,
|
||||
"url": INVALID_FIRMWARE_URL,
|
||||
}
|
||||
elif device.get("auto_update") is None:
|
||||
raise RuntimeError("auto_update is null")
|
||||
elif int(device["auto_update"]) != 0:
|
||||
ota_type = report.board.type if report.board is not None else None
|
||||
current_version = report.application.version if report.application is not None else None
|
||||
response["firmware"] = await self._firmware_info(ota_type, current_version, request_url)
|
||||
|
||||
websocket_url = await self.params.get_value("server.websocket", from_cache=True)
|
||||
auth_enabled = await self.params.get_value("server.auth.enabled", from_cache=True)
|
||||
websocket_token = ""
|
||||
if (auth_enabled or "").lower() == "true":
|
||||
try:
|
||||
websocket_token = await self._websocket_token(client_id, device_id)
|
||||
except Exception:
|
||||
logger.exception("WebSocket token generation failed")
|
||||
if is_blank(websocket_url) or websocket_url == "null":
|
||||
selected_websocket = "ws://xiaozhi.server.com:8000/xiaozhi/v1/"
|
||||
else:
|
||||
websocket_urls = _java_semicolon_split(websocket_url or "")
|
||||
selected_websocket = (
|
||||
random.choice(websocket_urls) # noqa: S311
|
||||
if websocket_urls
|
||||
else "ws://xiaozhi.server.com:8000/xiaozhi/v1/"
|
||||
)
|
||||
response["websocket"] = {"url": selected_websocket, "token": websocket_token}
|
||||
|
||||
mqtt_endpoint = await self.params.get_value("server.mqtt_gateway", from_cache=True)
|
||||
if mqtt_endpoint not in (None, "", "null"):
|
||||
try:
|
||||
group_id = str(device.get("board") or "GID_default") if device is not None else "GID_default"
|
||||
mqtt = await self._mqtt_config(device_id, group_id, client_ip)
|
||||
if mqtt is not None:
|
||||
mqtt["endpoint"] = mqtt_endpoint
|
||||
response["mqtt"] = mqtt
|
||||
except Exception:
|
||||
logger.exception("MQTT credential generation failed")
|
||||
|
||||
if device is None:
|
||||
response["activation"] = await self._activation(device_id, report)
|
||||
else:
|
||||
app_version = report.application.version if report.application is not None else None
|
||||
agent_id = device.get("agent_id")
|
||||
normalized_agent_id = None if agent_id is None else str(agent_id)
|
||||
if defer_connection_update is not None:
|
||||
defer_connection_update(str(device["id"]), normalized_agent_id, app_version)
|
||||
else:
|
||||
try:
|
||||
await self._persist_connection_update(
|
||||
str(device["id"]),
|
||||
normalized_agent_id,
|
||||
app_version,
|
||||
)
|
||||
except Exception:
|
||||
logger.exception("Asynchronous device connection update failed")
|
||||
return cast(dict[str, Any], self._drop_none(response))
|
||||
|
||||
async def _persist_connection_update(
|
||||
self,
|
||||
device_id: str,
|
||||
agent_id: str | None,
|
||||
app_version: str | None,
|
||||
) -> None:
|
||||
connection_time = shanghai_now_naive()
|
||||
try:
|
||||
await self.repository.update_connection(device_id, app_version=app_version, now=connection_time)
|
||||
await self.session.commit()
|
||||
except Exception:
|
||||
await self.session.rollback()
|
||||
raise
|
||||
if not is_blank(agent_id):
|
||||
await redis_set(f"agent:device:lastConnected:{agent_id}", connection_time, client=self.redis)
|
||||
|
||||
@staticmethod
|
||||
async def persist_connection_update_background(
|
||||
device_id: str,
|
||||
agent_id: str | None,
|
||||
app_version: str | None,
|
||||
) -> None:
|
||||
try:
|
||||
async with get_session_factory()() as session:
|
||||
await DeviceService(session)._persist_connection_update(device_id, agent_id, app_version)
|
||||
except Exception:
|
||||
logger.exception("Asynchronous device connection update failed")
|
||||
|
||||
async def ota_health_text(self) -> str:
|
||||
mqtt_gateway = await self.params.get_value("server.mqtt_gateway", from_cache=False)
|
||||
if is_blank(mqtt_gateway):
|
||||
return "OTA接口不正常,缺少mqtt_gateway地址,请登录智控台,在参数管理找到【server.mqtt_gateway】配置"
|
||||
websocket = await self.params.get_value("server.websocket", from_cache=True)
|
||||
if is_blank(websocket) or websocket == "null":
|
||||
return "OTA接口不正常,缺少websocket地址,请登录智控台,在参数管理找到【server.websocket】配置"
|
||||
ota_url = await self.params.get_value("server.ota", from_cache=True)
|
||||
if is_blank(ota_url) or ota_url == "null":
|
||||
return "OTA接口不正常,缺少ota地址,请登录智控台,在参数管理找到【server.ota】配置"
|
||||
return f"OTA接口运行正常,websocket集群数量:{len(_java_semicolon_split(websocket or ''))}"
|
||||
|
||||
async def ota_page(self, query: Mapping[str, Any]) -> dict[str, Any]:
|
||||
page = self._positive_int(query.get("page"), 1)
|
||||
limit = self._positive_int(query.get("limit"), 10)
|
||||
requested = query.get("orderField")
|
||||
requested_fields = [requested] if isinstance(requested, str) else list(requested or [])
|
||||
fields = [OTA_ORDER_COLUMNS[field] for field in requested_fields if field in OTA_ORDER_COLUMNS]
|
||||
if not fields:
|
||||
fields = ["update_date"]
|
||||
ascending = str(query.get("order") or "").lower() == "asc" if requested_fields else True
|
||||
firmware_name = query.get("firmwareName")
|
||||
name = str(firmware_name) if firmware_name is not None else None
|
||||
rows = await self.repository.list_ota(
|
||||
page=page,
|
||||
limit=limit,
|
||||
firmware_name=name,
|
||||
order_fields=fields,
|
||||
ascending=ascending,
|
||||
)
|
||||
rows = [self._ota_response_record(row) for row in rows]
|
||||
return {"total": await self.repository.count_ota(name), "list": rows}
|
||||
|
||||
async def get_ota_record(self, ota_id: str) -> dict[str, Any] | None:
|
||||
row = await self.repository.get_ota(ota_id)
|
||||
return None if row is None else self._ota_response_record(row)
|
||||
|
||||
async def save_ota(self, record: OtaRecord, user: AuthUser) -> None:
|
||||
values = record.model_dump(by_alias=False)
|
||||
existing = await self.repository.get_first_ota_by_type(record.type or "")
|
||||
now = shanghai_now_naive()
|
||||
if existing is not None:
|
||||
values["updater"] = record.updater if record.updater is not None else user.id
|
||||
values["update_date"] = record.update_date if record.update_date is not None else now
|
||||
await self.repository.update_ota(str(existing["id"]), values)
|
||||
else:
|
||||
values["id"] = record.id or uuid.uuid4().hex
|
||||
values["creator"] = record.creator if record.creator is not None else user.id
|
||||
values["updater"] = record.updater if record.updater is not None else user.id
|
||||
values["create_date"] = record.create_date if record.create_date is not None else now
|
||||
values["update_date"] = record.update_date if record.update_date is not None else now
|
||||
await self.repository.insert_ota(values)
|
||||
await self.session.commit()
|
||||
|
||||
async def update_ota(self, ota_id: str, record: OtaRecord, user: AuthUser) -> None:
|
||||
if await self.repository.count_duplicate_ota(
|
||||
ota_id=ota_id,
|
||||
ota_type=record.type,
|
||||
version=record.version,
|
||||
):
|
||||
raise RuntimeError("已存在相同类型和版本的固件,请修改后重试")
|
||||
values = record.model_dump(by_alias=False)
|
||||
values["updater"] = record.updater if record.updater is not None else user.id
|
||||
values["update_date"] = shanghai_now_naive()
|
||||
await self.repository.update_ota(ota_id, values)
|
||||
await self.session.commit()
|
||||
|
||||
async def delete_ota(self, ids: Sequence[str]) -> None:
|
||||
await self.repository.delete_ota(ids)
|
||||
await self.session.commit()
|
||||
|
||||
async def create_ota_download_id(self, ota_id: str) -> str:
|
||||
value = str(uuid.uuid4())
|
||||
await redis_set(f"ota:id:{value}", ota_id, client=self.redis)
|
||||
return value
|
||||
|
||||
async def resolve_ota_download(self, download_id: str) -> tuple[Path, str] | None:
|
||||
id_key = f"ota:id:{download_id}"
|
||||
ota_value = await redis_get(id_key, self.redis)
|
||||
if is_blank(None if ota_value is None else str(ota_value)):
|
||||
return None
|
||||
count_key = f"ota:download:count:{download_id}"
|
||||
count_value = await redis_get(count_key, self.redis)
|
||||
count = int(count_value or 0)
|
||||
if count >= 3:
|
||||
await redis_delete(count_key, id_key, client=self.redis)
|
||||
return None
|
||||
await redis_set(count_key, count + 1, client=self.redis)
|
||||
|
||||
ota_id = str(ota_value)
|
||||
if ota_id.startswith("file:"):
|
||||
firmware_path = ota_id[5:]
|
||||
ota_type = "assets"
|
||||
version = "1.0.0"
|
||||
else:
|
||||
record = await self.repository.get_ota(ota_id)
|
||||
firmware_value = None if record is None else record.get("firmware_path")
|
||||
if record is None or is_blank(None if firmware_value is None else str(firmware_value)):
|
||||
return None
|
||||
firmware_path = str(record["firmware_path"])
|
||||
ota_type = str(record.get("type"))
|
||||
version = str(record.get("version"))
|
||||
raw_path = Path(firmware_path)
|
||||
candidates = [raw_path] if raw_path.is_absolute() else [Path.cwd() / raw_path]
|
||||
if not raw_path.is_absolute() and raw_path.parts and raw_path.parts[0] == "uploadfile":
|
||||
candidates.insert(0, get_settings().upload_dir.joinpath(*raw_path.parts[1:]))
|
||||
candidates.append(Path.cwd() / "firmware" / raw_path.name)
|
||||
resolved = next((candidate for candidate in candidates if candidate.is_file()), None)
|
||||
if resolved is None:
|
||||
return None
|
||||
original_name = f"{ota_type}_{version}"
|
||||
dot_index = firmware_path.rfind(".")
|
||||
if dot_index >= 0:
|
||||
original_name += firmware_path[dot_index:]
|
||||
safe_name = re.sub(r"[^a-zA-Z0-9._-]", "_", original_name)
|
||||
return resolved, safe_name
|
||||
|
||||
async def save_firmware_file(self, *, filename: str | None, content: bytes) -> str:
|
||||
if not content:
|
||||
raise ValueError("上传文件不能为空")
|
||||
if filename is None:
|
||||
raise ValueError("文件名不能为空")
|
||||
dot_index = filename.rfind(".")
|
||||
if dot_index < 0:
|
||||
raise RuntimeError("文件名缺少扩展名")
|
||||
extension = filename[dot_index:].lower()
|
||||
if extension not in {".bin", ".apk"}:
|
||||
raise ValueError("只允许上传.bin和.apk格式的文件")
|
||||
digest = hashlib.md5(content, usedforsecurity=False).hexdigest()
|
||||
directory = get_settings().upload_dir
|
||||
directory.mkdir(parents=True, exist_ok=True)
|
||||
filename_on_disk = f"{digest}{extension}"
|
||||
physical_path = directory / filename_on_disk
|
||||
if not physical_path.exists():
|
||||
with physical_path.open("xb") as stream:
|
||||
stream.write(content)
|
||||
# Keep Java's database/API value stable even when the physical upload
|
||||
# volume is mounted elsewhere (for example /data/uploads in Docker).
|
||||
return str(Path("uploadfile") / filename_on_disk)
|
||||
|
||||
async def save_assets_file(self, *, filename: str | None, content: bytes, user: AuthUser) -> str:
|
||||
ota_url = await self.params.get_value("server.ota", from_cache=True)
|
||||
if is_blank(ota_url) or ota_url == "null":
|
||||
raise AppError(10102)
|
||||
if len(content) > 20 * 1024 * 1024:
|
||||
raise AppError(10142)
|
||||
if not user.is_super_admin:
|
||||
count_key = f"ota:upload:count:{user.id}"
|
||||
current = int(await redis_get(count_key, self.redis) or 0)
|
||||
if current >= 50:
|
||||
raise AppError(10195)
|
||||
await redis_increment(count_key, client=self.redis)
|
||||
path = await self.save_firmware_file(filename=filename, content=content)
|
||||
download_id = await self.create_ota_download_id(f"file:{path}")
|
||||
return (ota_url or "").replace("/ota/", "/otaMag/download/") + download_id
|
||||
|
||||
async def _gateway_device(self, device_id: str, user: AuthUser) -> tuple[str, dict[str, Any]] | None:
|
||||
gateway = await self.params.get_value("server.mqtt_manager_api", from_cache=True)
|
||||
if is_blank(gateway) or gateway == "null":
|
||||
return None
|
||||
device = await self.repository.get_device(device_id)
|
||||
if device is None or int(device.get("user_id") or -1) != user.id:
|
||||
return None
|
||||
return gateway or "", device
|
||||
|
||||
async def _post_gateway(self, url: str, body: Any, *, timeout_seconds: float | None = None) -> str:
|
||||
key = await self.params.get_value("server.mqtt_signature_key", from_cache=False)
|
||||
return await post_json(
|
||||
url,
|
||||
body,
|
||||
key or "",
|
||||
timeout_seconds=timeout_seconds or get_settings().external_request_timeout_seconds,
|
||||
client=self.http_client,
|
||||
)
|
||||
|
||||
async def _post_call(self, path: str, body: dict[str, Any], action: str) -> dict[str, Any]:
|
||||
gateway = await self.params.get_value("server.mqtt_manager_api", from_cache=True)
|
||||
key = await self.params.get_value("server.mqtt_signature_key", from_cache=True)
|
||||
if is_blank(gateway) or gateway == "null" or is_blank(key) or (key or "").strip().lower() == "null":
|
||||
return {"status": "error", "message": f"{action}失败,网关配置缺失"}
|
||||
result: dict[str, Any] = {"status": "error"}
|
||||
try:
|
||||
text = await post_json(
|
||||
f"http://{gateway}{path}",
|
||||
body,
|
||||
key or "",
|
||||
timeout_seconds=5.0,
|
||||
client=self.http_client,
|
||||
)
|
||||
if text.strip():
|
||||
payload = json.loads(text)
|
||||
if isinstance(payload, dict):
|
||||
result["status"] = payload.get("status")
|
||||
result["message"] = payload.get("message")
|
||||
return result
|
||||
except Exception:
|
||||
return {"status": "error", "message": f"{action}失败,请稍后再试"}
|
||||
|
||||
async def _firmware_info(
|
||||
self,
|
||||
ota_type: str | None,
|
||||
current_version: str | None,
|
||||
request_url: str,
|
||||
) -> dict[str, Any] | None:
|
||||
if is_blank(ota_type):
|
||||
return None
|
||||
selected_version = current_version if not is_blank(current_version) else "0.0.0"
|
||||
ota = await self.repository.get_latest_ota(ota_type or "")
|
||||
download_url: str | None = None
|
||||
if ota is not None and self._compare_versions(ota.get("version"), selected_version) > 0:
|
||||
ota_url = await self.params.get_value("server.ota", from_cache=True)
|
||||
if is_blank(ota_url) or ota_url == "null":
|
||||
ota_url = request_url
|
||||
download_id = await self.create_ota_download_id(str(ota["id"]))
|
||||
download_url = (ota_url or "").replace("/ota/", "/otaMag/download/") + download_id
|
||||
return {
|
||||
"version": selected_version if ota is None else ota.get("version"),
|
||||
"url": download_url or INVALID_FIRMWARE_URL,
|
||||
}
|
||||
|
||||
async def _activation(self, device_id: str, report: DeviceReportRequest) -> dict[str, Any]:
|
||||
safe_device_id = device_id.replace(":", "_").lower()
|
||||
data_key = f"ota:activation:data:{safe_device_id}"
|
||||
cached = _mapping(await redis_get(data_key, self.redis))
|
||||
code = str(cached.get("activation_code")) if cached and cached.get("activation_code") is not None else None
|
||||
frontend = await self.params.get_value("server.fronted_url", from_cache=True)
|
||||
if code is None or not code.strip():
|
||||
code = f"{secrets.randbelow(1_000_000):06d}"
|
||||
board = (
|
||||
report.board.type
|
||||
if report.board is not None and report.board.type is not None
|
||||
else (report.chip_model_name or "unknown")
|
||||
)
|
||||
app_version = report.application.version if report.application is not None else None
|
||||
await redis_set(
|
||||
data_key,
|
||||
{
|
||||
"id": device_id,
|
||||
"mac_address": device_id,
|
||||
"board": board,
|
||||
"app_version": app_version,
|
||||
"deviceId": device_id,
|
||||
"activation_code": code,
|
||||
},
|
||||
client=self.redis,
|
||||
)
|
||||
await redis_set(f"ota:activation:code:{code}", device_id, client=self.redis)
|
||||
return {
|
||||
"code": code,
|
||||
"message": f"{frontend if frontend is not None else 'null'}\n{code}",
|
||||
"challenge": device_id,
|
||||
}
|
||||
|
||||
async def _websocket_token(self, client_id: str, username: str) -> str:
|
||||
secret = await self.params.get_value("server.secret", from_cache=False)
|
||||
if is_blank(secret):
|
||||
raise RuntimeError("WebSocket认证密钥未配置(server.secret)")
|
||||
timestamp = int(datetime.now().timestamp())
|
||||
message = f"{client_id}|{username}|{timestamp}".encode()
|
||||
signature = hmac.new((secret or "").encode(), message, hashlib.sha256).digest()
|
||||
encoded = base64.urlsafe_b64encode(signature).decode().rstrip("=")
|
||||
return f"{encoded}.{timestamp}"
|
||||
|
||||
async def _mqtt_config(self, mac_address: str, group_id: str, client_ip: str) -> dict[str, Any] | None:
|
||||
key = await self.params.get_value("server.mqtt_signature_key", from_cache=True)
|
||||
if is_blank(key):
|
||||
return None
|
||||
client_id = self._mqtt_client_id(group_id, mac_address)
|
||||
user_data = json.dumps({"ip": client_ip}, ensure_ascii=False, separators=(",", ":"))
|
||||
username = base64.b64encode(user_data.encode()).decode()
|
||||
password = base64.b64encode(
|
||||
hmac.new((key or "").encode(), f"{client_id}|{username}".encode(), hashlib.sha256).digest()
|
||||
).decode()
|
||||
safe_mac = mac_address.replace(":", "_")
|
||||
return {
|
||||
"client_id": client_id,
|
||||
"username": username,
|
||||
"password": password,
|
||||
"publish_topic": "device-server",
|
||||
"subscribe_topic": f"devices/p2p/{safe_mac}",
|
||||
}
|
||||
|
||||
@staticmethod
|
||||
def _mqtt_client_id(group_id: str, mac_address: str) -> str:
|
||||
safe_group = group_id.replace(":", "_")
|
||||
safe_mac = mac_address.replace(":", "_")
|
||||
return f"{safe_group}@@@{safe_mac}@@@{safe_mac}"
|
||||
|
||||
@staticmethod
|
||||
def _compare_versions(first: Any, second: Any) -> int:
|
||||
if first is None or second is None:
|
||||
return 0
|
||||
first = str(first)
|
||||
second = str(second)
|
||||
first_parts = first.split(".")
|
||||
second_parts = second.split(".")
|
||||
for index in range(max(len(first_parts), len(second_parts))):
|
||||
first_value = int(first_parts[index]) if index < len(first_parts) else 0
|
||||
second_value = int(second_parts[index]) if index < len(second_parts) else 0
|
||||
if first_value != second_value:
|
||||
return 1 if first_value > second_value else -1
|
||||
return 0
|
||||
|
||||
async def _unique_alias(self, mac_address: str, alias: str | None) -> str | None:
|
||||
existing = await self.repository.get_aliases(mac_address)
|
||||
if alias not in existing:
|
||||
return alias
|
||||
suffix = 1
|
||||
while f"{alias}{suffix}" in existing:
|
||||
suffix += 1
|
||||
return f"{alias}{suffix}"
|
||||
|
||||
@staticmethod
|
||||
def _mac_device_name(mac: str) -> str:
|
||||
return mac if len(mac) < 2 else f"尾号为{mac[-2:]}的设备"
|
||||
|
||||
@staticmethod
|
||||
def _positive_int(value: Any, default: int) -> int:
|
||||
if value is None:
|
||||
return default
|
||||
return int(str(value))
|
||||
|
||||
@staticmethod
|
||||
def _drop_none(value: Any) -> Any:
|
||||
if isinstance(value, dict):
|
||||
return {key: DeviceService._drop_none(item) for key, item in value.items() if item is not None}
|
||||
if isinstance(value, list):
|
||||
return [DeviceService._drop_none(item) for item in value]
|
||||
return value
|
||||
|
||||
@staticmethod
|
||||
def _user_device_view(row: Mapping[str, Any]) -> dict[str, Any]:
|
||||
return {
|
||||
"app_version": row.get("app_version"),
|
||||
"bind_user_name": None,
|
||||
"device_type": row.get("board"),
|
||||
"board": row.get("board"),
|
||||
"id": row.get("id"),
|
||||
"mac_address": row.get("mac_address"),
|
||||
"alias": row.get("alias"),
|
||||
"ota_upgrade": None,
|
||||
"recent_chat_time": None,
|
||||
"last_connected_at_timestamp": DeviceService._timestamp(row.get("last_connected_at")),
|
||||
"create_date_timestamp": DeviceService._timestamp(row.get("create_date")),
|
||||
# UserShowDeviceListVO pins only this field to UTC. The companion
|
||||
# epoch value still uses the configured Asia/Shanghai instant.
|
||||
"create_date": DeviceService._utc_datetime(row.get("create_date")),
|
||||
}
|
||||
|
||||
@staticmethod
|
||||
def _utc_datetime(value: Any) -> Any:
|
||||
if not isinstance(value, datetime):
|
||||
return value
|
||||
localized = value.replace(tzinfo=ZoneInfo(get_settings().timezone)) if value.tzinfo is None else value
|
||||
return localized.astimezone(timezone.utc).replace(tzinfo=None)
|
||||
|
||||
@staticmethod
|
||||
def _ota_response_record(row: Mapping[str, Any]) -> dict[str, Any]:
|
||||
result = dict(row)
|
||||
if result.get("size") is not None:
|
||||
result["size"] = str(result["size"])
|
||||
return result
|
||||
|
||||
@staticmethod
|
||||
def _timestamp(value: Any) -> int | None:
|
||||
if not isinstance(value, datetime):
|
||||
return None
|
||||
localized = value.replace(tzinfo=ZoneInfo(get_settings().timezone)) if value.tzinfo is None else value
|
||||
return int(localized.timestamp() * 1000)
|
||||
@@ -0,0 +1,21 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from functools import lru_cache
|
||||
from pathlib import Path
|
||||
|
||||
from app.core.config import get_settings
|
||||
from app.core.i18n import LANGUAGE_FILES, _load_properties, resolve_language
|
||||
|
||||
|
||||
@lru_cache(maxsize=16)
|
||||
def _validation_messages(language: str, directory: str) -> dict[str, str]:
|
||||
root = Path(directory)
|
||||
values = _load_properties(root / "validation.properties")
|
||||
localized = LANGUAGE_FILES[language].replace("messages_", "validation_")
|
||||
values.update(_load_properties(root / localized))
|
||||
return values
|
||||
|
||||
|
||||
def validation_message(key: str, accept_language: str | None) -> str:
|
||||
language = resolve_language(accept_language)
|
||||
return _validation_messages(language, str(get_settings().i18n_dir)).get(key, key)
|
||||
@@ -0,0 +1,723 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
from collections import defaultdict
|
||||
from datetime import datetime
|
||||
from typing import Any
|
||||
from zoneinfo import ZoneInfo
|
||||
|
||||
from fastapi import UploadFile
|
||||
|
||||
from app.core.config import get_settings
|
||||
from app.core.errors import AppError
|
||||
from app.core.i18n import message_for
|
||||
from app.core.redis import get_redis
|
||||
from app.core.security import AuthUser, shanghai_now_naive
|
||||
from app.core.serialization import preserve_java_map_keys
|
||||
from app.integrations.ragflow import RAGFlowClient
|
||||
from app.repositories.knowledge import KnowledgeRepository
|
||||
from app.schemas.knowledge import KnowledgeBaseBody, RetrievalBody
|
||||
|
||||
|
||||
def dataset_dto(row: dict[str, Any]) -> dict[str, Any]:
|
||||
return {
|
||||
"id": row.get("id"),
|
||||
"datasetId": row.get("dataset_id"),
|
||||
"ragModelId": row.get("rag_model_id"),
|
||||
"name": row.get("name"),
|
||||
"avatar": row.get("avatar"),
|
||||
"description": row.get("description"),
|
||||
"embeddingModel": row.get("embedding_model"),
|
||||
"permission": row.get("permission"),
|
||||
"chunkMethod": row.get("chunk_method"),
|
||||
"parserConfig": row.get("parser_config"),
|
||||
"chunkCount": None if row.get("chunk_count") is None else str(row["chunk_count"]),
|
||||
"tokenNum": None if row.get("token_num") is None else str(row["token_num"]),
|
||||
"status": row.get("status"),
|
||||
"creator": row.get("creator"),
|
||||
"createdAt": row.get("created_at"),
|
||||
"updater": row.get("updater"),
|
||||
"updatedAt": row.get("updated_at"),
|
||||
# KnowledgeBaseEntity.documentCount is Long while KnowledgeBaseDTO uses
|
||||
# Integer. Spring BeanUtils does not coerce that property, so local DTO
|
||||
# conversion leaves it null; list enrichment fills it from RAGFlow.
|
||||
"documentCount": None,
|
||||
"errorMessage": row.get("error_message"),
|
||||
}
|
||||
|
||||
|
||||
def document_dto(row: dict[str, Any]) -> dict[str, Any]:
|
||||
return {
|
||||
"id": row.get("document_id"),
|
||||
"documentId": row.get("document_id"),
|
||||
"datasetId": row.get("dataset_id"),
|
||||
"name": row.get("name"),
|
||||
# RAGFlowAdapter.mapToKnowledgeFilesDTO does not populate these two
|
||||
# fields for the immediate upload response.
|
||||
"fileType": None,
|
||||
"fileSize": row.get("size"),
|
||||
"filePath": None,
|
||||
"progress": row.get("progress"),
|
||||
"thumbnail": row.get("thumbnail"),
|
||||
"processDuration": row.get("process_duration"),
|
||||
"sourceType": row.get("source_type"),
|
||||
"metaFields": _json_object(row.get("meta_fields")),
|
||||
"chunkMethod": row.get("chunk_method"),
|
||||
"parserConfig": _json_object(row.get("parser_config")),
|
||||
"status": row.get("status"),
|
||||
"run": row.get("run"),
|
||||
"creator": row.get("creator"),
|
||||
"createdAt": row.get("created_at"),
|
||||
"updater": None,
|
||||
"updatedAt": row.get("updated_at"),
|
||||
"chunkCount": row.get("chunk_count"),
|
||||
"tokenCount": row.get("token_count"),
|
||||
"error": row.get("error"),
|
||||
"parseStatusCode": _parse_status(row.get("run")),
|
||||
}
|
||||
|
||||
|
||||
def remote_document_dto(row: dict[str, Any], dataset_id: str) -> dict[str, Any]:
|
||||
run = row.get("run")
|
||||
return {
|
||||
"id": row.get("id"),
|
||||
"documentId": row.get("id"),
|
||||
"datasetId": row.get("dataset_id") or dataset_id,
|
||||
"name": row.get("name"),
|
||||
"fileType": row.get("type"),
|
||||
"fileSize": row.get("size"),
|
||||
"filePath": None,
|
||||
"progress": row.get("progress"),
|
||||
"thumbnail": row.get("thumbnail"),
|
||||
"processDuration": row.get("process_duration"),
|
||||
"sourceType": row.get("source_type"),
|
||||
"metaFields": row.get("meta_fields"),
|
||||
"chunkMethod": row.get("chunk_method"),
|
||||
"parserConfig": row.get("parser_config"),
|
||||
"status": _remote_status(row.get("status")),
|
||||
"run": run,
|
||||
"creator": None,
|
||||
"createdAt": _millis(row.get("create_time")),
|
||||
"updater": None,
|
||||
"updatedAt": _millis(row.get("update_time")),
|
||||
"chunkCount": row.get("chunk_count") or 0,
|
||||
"tokenCount": row.get("token_count"),
|
||||
"error": row.get("progress_msg"),
|
||||
"parseStatusCode": _parse_status(run),
|
||||
}
|
||||
|
||||
|
||||
def _parse_status(run: Any) -> int:
|
||||
return {"RUNNING": 1, "CANCEL": 2, "DONE": 3, "FAIL": 4}.get(str(run or "").upper(), 0)
|
||||
|
||||
|
||||
def _json_object(value: Any) -> dict[str, Any] | None:
|
||||
if value is None:
|
||||
return None
|
||||
if isinstance(value, dict):
|
||||
return dict(value)
|
||||
try:
|
||||
parsed = json.loads(value.decode() if isinstance(value, bytes) else str(value))
|
||||
return dict(parsed) if isinstance(parsed, dict) else None
|
||||
except (ValueError, TypeError):
|
||||
return None
|
||||
|
||||
|
||||
def _millis(value: Any) -> Any:
|
||||
try:
|
||||
if value is None:
|
||||
return None
|
||||
timezone = ZoneInfo(get_settings().timezone)
|
||||
return datetime.fromtimestamp(float(value) / 1000, timezone).replace(tzinfo=None)
|
||||
except (TypeError, ValueError, OSError):
|
||||
return None
|
||||
|
||||
|
||||
def _is_blank(value: str | None) -> bool:
|
||||
return value is None or not value.strip()
|
||||
|
||||
|
||||
def _remote_status(value: Any) -> str:
|
||||
if value is None or (isinstance(value, str) and not value.strip()):
|
||||
return "1"
|
||||
return str(value)
|
||||
|
||||
|
||||
class KnowledgeBaseService:
|
||||
def __init__(self, repository: KnowledgeRepository):
|
||||
self.repository = repository
|
||||
|
||||
async def get_owned(self, identifier: str, user: AuthUser) -> dict[str, Any]:
|
||||
if not identifier.strip():
|
||||
raise AppError(10003)
|
||||
row = await self.repository.get_dataset(identifier)
|
||||
if row is None:
|
||||
raise AppError(10163)
|
||||
if row.get("creator") is None or int(row["creator"]) != user.id:
|
||||
raise AppError(10169)
|
||||
return row
|
||||
|
||||
async def page(
|
||||
self,
|
||||
user: AuthUser,
|
||||
name: str | None,
|
||||
page: int,
|
||||
page_size: int,
|
||||
language: str | None = None,
|
||||
) -> dict[str, Any]:
|
||||
rows, total = await self.repository.dataset_page(
|
||||
user.id, name, (max(page, 1) - 1) * page_size, page_size
|
||||
)
|
||||
results: list[dict[str, Any]] = []
|
||||
changed = False
|
||||
for row in rows:
|
||||
dto = dataset_dto(row)
|
||||
if row.get("dataset_id") and row.get("rag_model_id"):
|
||||
try:
|
||||
client = await self._client(str(row["rag_model_id"]))
|
||||
remote = await client.dataset_info(str(row["dataset_id"]))
|
||||
if remote is None:
|
||||
await self.repository.execute(
|
||||
"DELETE FROM ai_rag_knowledge_document WHERE dataset_id=:dataset_id",
|
||||
{"dataset_id": row["dataset_id"]},
|
||||
)
|
||||
await self.repository.delete_dataset_local(row)
|
||||
await _delete_cache_ignoring_errors(f"knowledge:base:{row['id']}")
|
||||
changed = True
|
||||
continue
|
||||
remote_name = remote.get("name")
|
||||
local_name = (
|
||||
str(remote_name).split("_", 1)[1]
|
||||
if remote_name and "_" in str(remote_name)
|
||||
else remote_name
|
||||
)
|
||||
updates: dict[str, Any] = {}
|
||||
if local_name and local_name != row.get("name"):
|
||||
updates["name"] = local_name
|
||||
dto["name"] = local_name
|
||||
if remote.get("description") != row.get("description"):
|
||||
updates["description"] = remote.get("description")
|
||||
dto["description"] = remote.get("description")
|
||||
if updates:
|
||||
await self.repository.execute(
|
||||
"UPDATE ai_rag_dataset SET name=COALESCE(:name,name),description=:description WHERE id=:id",
|
||||
{
|
||||
"name": updates.get("name"),
|
||||
"description": updates.get("description", row.get("description")),
|
||||
"id": row["id"],
|
||||
},
|
||||
)
|
||||
changed = True
|
||||
if remote.get("document_count") is not None:
|
||||
dto["documentCount"] = int(remote["document_count"])
|
||||
except Exception as exc:
|
||||
dto["documentCount"] = 0
|
||||
dto["errorMessage"] = (
|
||||
message_for(exc.code, language, *exc.params)
|
||||
if isinstance(exc, AppError)
|
||||
else str(exc)
|
||||
)
|
||||
results.append(dto)
|
||||
if changed:
|
||||
await self.repository.session.commit()
|
||||
return {"total": total, "list": results}
|
||||
|
||||
async def create(self, body: KnowledgeBaseBody, user: AuthUser) -> dict[str, Any]:
|
||||
if not _is_blank(body.name) and await self.repository.duplicate_dataset_name(user.id, str(body.name)):
|
||||
raise AppError(10170)
|
||||
rag_model_id = body.rag_model_id
|
||||
if _is_blank(rag_model_id):
|
||||
models = await self.repository.rag_models()
|
||||
if not models:
|
||||
raise AppError(10164, params=("未指定且无可用默认 RAG 模型",))
|
||||
rag_model_id = str(models[0]["id"])
|
||||
client = await self._client(str(rag_model_id))
|
||||
create_body = {
|
||||
"name": f"{user.username}_{'null' if body.name is None else body.name}",
|
||||
"avatar": body.avatar,
|
||||
"description": body.description,
|
||||
"embedding_model": body.embedding_model,
|
||||
"permission": body.permission,
|
||||
"chunk_method": body.chunk_method,
|
||||
# KnowledgeBaseDTO.parserConfig is a String, while CreateReq uses
|
||||
# ParserConfig. BeanUtils skips the incompatible property.
|
||||
"parser_config": None,
|
||||
}
|
||||
remote = await client.create_dataset(create_body)
|
||||
dataset_id = str(remote["id"])
|
||||
now = shanghai_now_naive()
|
||||
created_at = body.created_at or now
|
||||
updated_at = body.updated_at or now
|
||||
values = {
|
||||
"id": dataset_id,
|
||||
"dataset_id": dataset_id,
|
||||
"rag_model_id": rag_model_id,
|
||||
"tenant_id": remote.get("tenant_id"),
|
||||
"name": body.name,
|
||||
"avatar": remote.get("avatar") if _is_blank(body.avatar) else body.avatar,
|
||||
"description": body.description,
|
||||
"embedding_model": remote.get("embedding_model"),
|
||||
"permission": remote.get("permission"),
|
||||
"chunk_method": remote.get("chunk_method"),
|
||||
"parser_config": json.dumps(
|
||||
remote.get("parser_config"), ensure_ascii=False, separators=(",", ":")
|
||||
)
|
||||
if remote.get("parser_config") is not None
|
||||
else None,
|
||||
"chunk_count": remote.get("chunk_count") or 0,
|
||||
"document_count": remote.get("document_count") or 0,
|
||||
"token_num": remote.get("token_num") or 0,
|
||||
"status": 1,
|
||||
"creator": user.id,
|
||||
"updater": user.id,
|
||||
"created_at": created_at,
|
||||
"updated_at": updated_at,
|
||||
}
|
||||
try:
|
||||
await self.repository.insert_dataset(values)
|
||||
await self.repository.session.commit()
|
||||
except Exception as exc:
|
||||
await self.repository.session.rollback()
|
||||
try:
|
||||
await client.delete_datasets([dataset_id])
|
||||
except AppError:
|
||||
pass
|
||||
if isinstance(exc, AppError):
|
||||
raise
|
||||
raise AppError(10167, params=(f"创建知识库失败: {exc}",)) from exc
|
||||
return dataset_dto(values)
|
||||
|
||||
async def update(
|
||||
self, identifier: str, body: KnowledgeBaseBody, user: AuthUser
|
||||
) -> dict[str, Any]:
|
||||
existing = await self.get_owned(identifier, user)
|
||||
if not _is_blank(body.name) and await self.repository.duplicate_dataset_name(
|
||||
user.id, str(body.name), str(existing["id"])
|
||||
):
|
||||
raise AppError(10170)
|
||||
if not _is_blank(identifier) and await self.repository.dataset_id_conflict(
|
||||
identifier, str(existing["id"])
|
||||
):
|
||||
raise AppError(10002)
|
||||
rag_model_id = body.rag_model_id
|
||||
effective_permission = body.permission
|
||||
effective_chunk_method = body.chunk_method
|
||||
if existing.get("dataset_id") and not _is_blank(rag_model_id):
|
||||
if _is_blank(effective_permission):
|
||||
effective_permission = existing.get("permission")
|
||||
if _is_blank(effective_chunk_method):
|
||||
effective_chunk_method = existing.get("chunk_method")
|
||||
client = await self._client(str(rag_model_id))
|
||||
remote_body = {
|
||||
"name": f"{user.username}_{body.name}" if not _is_blank(body.name) else None,
|
||||
"avatar": body.avatar,
|
||||
"description": body.description,
|
||||
"embedding_model": body.embedding_model,
|
||||
"permission": effective_permission,
|
||||
"chunk_method": effective_chunk_method,
|
||||
"parser_config": _json_object(body.parser_config),
|
||||
}
|
||||
await client.update_dataset(str(existing["dataset_id"]), remote_body)
|
||||
now = shanghai_now_naive()
|
||||
updater = body.updater if body.updater is not None else user.id
|
||||
updated_at = body.updated_at or now
|
||||
values = {
|
||||
"id": existing["id"],
|
||||
# The controller injects the literal path value into datasetId,
|
||||
# even when a legacy row was found through its local primary key.
|
||||
"dataset_id": identifier,
|
||||
"rag_model_id": rag_model_id,
|
||||
"name": body.name,
|
||||
"avatar": body.avatar,
|
||||
"description": body.description,
|
||||
"embedding_model": body.embedding_model,
|
||||
"permission": effective_permission,
|
||||
"chunk_method": effective_chunk_method,
|
||||
"parser_config": body.parser_config,
|
||||
"chunk_count": body.chunk_count,
|
||||
"token_num": body.token_num,
|
||||
"status": body.status,
|
||||
"creator": body.creator,
|
||||
"created_at": body.created_at,
|
||||
"updater": updater,
|
||||
"updated_at": updated_at,
|
||||
}
|
||||
try:
|
||||
await self.repository.update_dataset(values)
|
||||
# Java performs cache eviction inside the database transaction;
|
||||
# an eviction failure therefore rolls this update back.
|
||||
await get_redis().delete(f"knowledge:base:{existing['id']}")
|
||||
await self.repository.session.commit()
|
||||
except Exception:
|
||||
await self.repository.session.rollback()
|
||||
raise
|
||||
# BeanUtils copies request nulls onto the in-memory entity before
|
||||
# MyBatis' NOT_NULL update strategy preserves the stored columns. The
|
||||
# Java response is built from that in-memory entity, so its null fields
|
||||
# intentionally differ from a subsequent GET of the row.
|
||||
return dataset_dto(values)
|
||||
|
||||
async def delete(self, identifier: str, user: AuthUser, language: str | None = None) -> None:
|
||||
row = await self.get_owned(identifier, user)
|
||||
documents = await self.repository.all_documents(str(row["dataset_id"]))
|
||||
if documents:
|
||||
# Java's document orchestration necessarily resolves the adapter
|
||||
# when child records exist.
|
||||
client = await self._client(str(row.get("rag_model_id") or ""))
|
||||
ids = [str(item["document_id"]) for item in documents]
|
||||
if any(item.get("run") == "RUNNING" for item in documents):
|
||||
raise AppError(10199)
|
||||
try:
|
||||
await client.delete_documents(str(row["dataset_id"]), ids)
|
||||
except Exception as exc:
|
||||
raise _document_delete_error(exc, language) from exc
|
||||
await self.repository.delete_document_shadows(str(row["dataset_id"]), ids)
|
||||
await self.repository.update_stats(
|
||||
str(row["dataset_id"]),
|
||||
-len(ids),
|
||||
-sum(int(item.get("chunk_count") or 0) for item in documents),
|
||||
-sum(int(item.get("token_count") or 0) for item in documents),
|
||||
)
|
||||
# deleteDocuments is NOT_SUPPORTED in Java and its shadow cleanup
|
||||
# commits before the outer dataset transaction continues.
|
||||
await self.repository.session.commit()
|
||||
await _delete_cache_ignoring_errors(f"knowledge:base:{row['dataset_id']}")
|
||||
if not _is_blank(row.get("rag_model_id")) and not _is_blank(row.get("dataset_id")):
|
||||
client = await self._client(str(row["rag_model_id"]))
|
||||
await client.delete_datasets([str(row["dataset_id"])])
|
||||
await self.repository.delete_dataset_local(row)
|
||||
try:
|
||||
await get_redis().delete(f"knowledge:base:{row['id']}")
|
||||
await self.repository.session.commit()
|
||||
except Exception:
|
||||
await self.repository.session.rollback()
|
||||
raise
|
||||
|
||||
async def batch_delete(
|
||||
self, identifiers: list[str], user: AuthUser, language: str | None = None
|
||||
) -> None:
|
||||
rows = await self.repository.datasets_by_ids(identifiers)
|
||||
for row in rows:
|
||||
if row.get("creator") is None or int(row["creator"]) != user.id:
|
||||
raise AppError(10169)
|
||||
# Preserve Java's sequential external calls and stop-on-first-error semantics.
|
||||
for row in rows:
|
||||
await self.delete(str(row["dataset_id"]), user, language)
|
||||
|
||||
async def rag_models(self) -> list[dict[str, Any]]:
|
||||
rows = await self.repository.rag_models()
|
||||
result: list[dict[str, Any]] = []
|
||||
for row in rows:
|
||||
result.append(
|
||||
{
|
||||
"id": row.get("id"),
|
||||
"modelType": None,
|
||||
"modelCode": None,
|
||||
"modelName": row.get("model_name"),
|
||||
"isDefault": None,
|
||||
"isEnabled": None,
|
||||
# ModelConfigEntity.configJson is a JSONObject. Jackson
|
||||
# preserves its dynamic snake_case keys instead of applying
|
||||
# the DTO property naming strategy recursively.
|
||||
"configJson": preserve_java_map_keys(_json_object(row.get("config_json"))),
|
||||
"docLink": None,
|
||||
"remark": None,
|
||||
"sort": None,
|
||||
"updater": None,
|
||||
"updateDate": None,
|
||||
"creator": None,
|
||||
"createDate": None,
|
||||
}
|
||||
)
|
||||
return result
|
||||
|
||||
async def _client(self, model_id: str) -> RAGFlowClient:
|
||||
config = await self.repository.rag_config(model_id)
|
||||
adapter_type = config.get("type")
|
||||
if adapter_type != "ragflow":
|
||||
raise AppError(10184, params=(f"适配器类型未注册: {adapter_type}",))
|
||||
try:
|
||||
return RAGFlowClient(config)
|
||||
except AppError as exc:
|
||||
# KnowledgeBaseAdapterFactory wraps adapter initialization and
|
||||
# validateConfig failures as RAG_ADAPTER_CREATION_FAILED.
|
||||
if exc.code in {10171, 10172, 10173, 10174}:
|
||||
raise AppError(10186) from exc
|
||||
raise
|
||||
|
||||
|
||||
class KnowledgeDocumentService:
|
||||
def __init__(self, repository: KnowledgeRepository):
|
||||
self.repository = repository
|
||||
self.datasets = KnowledgeBaseService(repository)
|
||||
|
||||
async def page(
|
||||
self,
|
||||
dataset_id: str,
|
||||
user: AuthUser,
|
||||
*,
|
||||
name: str | None,
|
||||
status: str | None,
|
||||
page: int,
|
||||
page_size: int,
|
||||
) -> dict[str, Any]:
|
||||
await self.datasets.get_owned(dataset_id, user)
|
||||
try:
|
||||
await self.reconcile(dataset_id, creator=user.id)
|
||||
except Exception:
|
||||
await self.repository.session.rollback()
|
||||
rows, total = await self.repository.documents_page(
|
||||
dataset_id,
|
||||
name=name,
|
||||
status=status,
|
||||
offset=(max(page, 1) - 1) * page_size,
|
||||
limit=page_size,
|
||||
)
|
||||
return {"total": total, "list": [document_dto(row) for row in rows]}
|
||||
|
||||
async def upload(
|
||||
self,
|
||||
dataset_id: str,
|
||||
user: AuthUser,
|
||||
file: UploadFile,
|
||||
*,
|
||||
name: str | None,
|
||||
meta_fields: dict[str, Any] | None,
|
||||
chunk_method: str | None,
|
||||
parser_config: dict[str, Any] | None,
|
||||
) -> dict[str, Any]:
|
||||
await self.datasets.get_owned(dataset_id, user)
|
||||
content = await file.read()
|
||||
if not dataset_id.strip() or not content:
|
||||
raise AppError(10003)
|
||||
file_name = file.filename if _is_blank(name) else name
|
||||
if _is_blank(file_name):
|
||||
raise AppError(10179)
|
||||
assert file_name is not None
|
||||
client = await self._client_for_dataset(dataset_id)
|
||||
remote = await client.upload_document(
|
||||
dataset_id,
|
||||
file,
|
||||
content,
|
||||
name=file_name,
|
||||
meta_fields=meta_fields,
|
||||
chunk_method=chunk_method,
|
||||
parser_config=parser_config,
|
||||
)
|
||||
if not remote.get("id"):
|
||||
raise AppError(10167, params=("远程上传成功但未返回有效 DocumentID",))
|
||||
remote.setdefault("dataset_id", dataset_id)
|
||||
shadow = dict(remote)
|
||||
if _is_blank(str(shadow.get("name")) if shadow.get("name") is not None else None):
|
||||
shadow["name"] = file_name
|
||||
# Java stores the original controller values in the shadow row, even
|
||||
# when invalid chunk methods were omitted from the RAGFlow request.
|
||||
shadow["chunk_method"] = chunk_method
|
||||
shadow["parser_config"] = parser_config
|
||||
inserted = await self.repository.upsert_document(dataset_id, shadow, creator=user.id)
|
||||
if inserted:
|
||||
await self.repository.update_stats(dataset_id, 1, 0, 0)
|
||||
await self.repository.session.commit()
|
||||
return remote_document_dto(remote, dataset_id)
|
||||
|
||||
async def delete(
|
||||
self,
|
||||
dataset_id: str,
|
||||
ids: list[str] | None,
|
||||
user: AuthUser,
|
||||
language: str | None = None,
|
||||
) -> None:
|
||||
await self.datasets.get_owned(dataset_id, user)
|
||||
if not ids:
|
||||
raise AppError(10178)
|
||||
rows = await self.repository.documents_by_remote_ids(dataset_id, ids)
|
||||
if len(rows) != len(ids):
|
||||
raise AppError(10169)
|
||||
if any(row.get("run") == "RUNNING" for row in rows):
|
||||
raise AppError(10199)
|
||||
chunks = sum(int(row.get("chunk_count") or 0) for row in rows)
|
||||
tokens = sum(int(row.get("token_count") or 0) for row in rows)
|
||||
client = await self._client_for_dataset(dataset_id)
|
||||
try:
|
||||
await client.delete_documents(dataset_id, ids)
|
||||
except Exception as exc:
|
||||
raise _document_delete_error(exc, language) from exc
|
||||
deleted = await self.repository.delete_document_shadows(dataset_id, ids)
|
||||
if deleted:
|
||||
await self.repository.update_stats(dataset_id, -len(ids), -chunks, -tokens)
|
||||
await self.repository.session.commit()
|
||||
await _delete_cache_ignoring_errors(f"knowledge:base:{dataset_id}")
|
||||
|
||||
async def parse(self, dataset_id: str, ids: list[str], user: AuthUser) -> bool:
|
||||
await self.datasets.get_owned(dataset_id, user)
|
||||
if not ids:
|
||||
raise AppError(10178)
|
||||
client = await self._client_for_dataset(dataset_id)
|
||||
await client.parse_documents(dataset_id, ids)
|
||||
await self.repository.mark_documents_running(dataset_id, ids, shanghai_now_naive())
|
||||
await self.repository.session.commit()
|
||||
return True
|
||||
|
||||
async def chunks(
|
||||
self,
|
||||
dataset_id: str,
|
||||
document_id: str,
|
||||
user: AuthUser,
|
||||
*,
|
||||
page: int,
|
||||
page_size: int,
|
||||
keywords: str | None,
|
||||
chunk_id: str | None,
|
||||
) -> dict[str, Any]:
|
||||
await self.datasets.get_owned(dataset_id, user)
|
||||
client = await self._client_for_dataset(dataset_id)
|
||||
return await client.chunks(
|
||||
dataset_id,
|
||||
document_id,
|
||||
{"page": page, "page_size": page_size, "keywords": keywords, "id": chunk_id},
|
||||
)
|
||||
|
||||
async def retrieval(self, dataset_id: str, body: RetrievalBody, user: AuthUser) -> dict[str, Any]:
|
||||
await self.datasets.get_owned(dataset_id, user)
|
||||
dataset_ids = body.dataset_ids or [dataset_id]
|
||||
if not dataset_ids:
|
||||
raise AppError(500, "未指定召回测试的知识库")
|
||||
page = body.page if body.page is not None and body.page >= 1 else 1
|
||||
page_size = body.page_size if body.page_size is not None and body.page_size >= 1 else 100
|
||||
top_k = body.top_k if body.top_k is None or body.top_k >= 1 else 1024
|
||||
threshold = body.similarity_threshold
|
||||
if threshold is not None:
|
||||
threshold = 0.2 if threshold < 0 else min(threshold, 1.0)
|
||||
payload: dict[str, Any] = {
|
||||
"dataset_ids": dataset_ids,
|
||||
"document_ids": body.document_ids,
|
||||
"question": body.question,
|
||||
"page": page,
|
||||
"page_size": page_size,
|
||||
"similarity_threshold": threshold,
|
||||
"vector_similarity_weight": body.vector_similarity_weight,
|
||||
"top_k": top_k,
|
||||
"rerank_id": body.rerank_id,
|
||||
"highlight": body.highlight,
|
||||
"keyword": body.keyword,
|
||||
"cross_languages": body.cross_languages,
|
||||
"metadata_condition": body.metadata_condition,
|
||||
}
|
||||
payload = {key: value for key, value in payload.items() if value is not None}
|
||||
client = await self._client_for_dataset(dataset_ids[0])
|
||||
return await client.retrieval(payload)
|
||||
|
||||
async def reconcile(self, dataset_id: str, *, creator: int | None = None) -> int:
|
||||
client = await self._client_for_dataset(dataset_id)
|
||||
remote: list[dict[str, Any]] = []
|
||||
page, total = 1, 2**63 - 1
|
||||
while (page - 1) * 100 < total:
|
||||
rows, total = await client.documents(dataset_id, page=page, page_size=100)
|
||||
if not rows:
|
||||
break
|
||||
remote.extend(rows)
|
||||
page += 1
|
||||
local = await self.repository.all_documents(dataset_id)
|
||||
remote_map = {str(item.get("id")): item for item in remote if item.get("id")}
|
||||
local_map = {str(item["document_id"]): item for item in local}
|
||||
new_count = 0
|
||||
for document_id, item in remote_map.items():
|
||||
prior = local_map.get(document_id)
|
||||
inserted = await self.repository.upsert_document(dataset_id, item, creator=creator)
|
||||
if inserted:
|
||||
new_count += 1
|
||||
await self.repository.update_stats(
|
||||
dataset_id, 1, int(item.get("chunk_count") or 0), int(item.get("token_count") or 0)
|
||||
)
|
||||
elif prior:
|
||||
await self.repository.update_stats(
|
||||
dataset_id,
|
||||
0,
|
||||
int(item.get("chunk_count") or 0) - int(prior.get("chunk_count") or 0),
|
||||
int(item.get("token_count") or 0) - int(prior.get("token_count") or 0),
|
||||
)
|
||||
deleted_ids = [identifier for identifier in local_map if identifier not in remote_map]
|
||||
if deleted_ids:
|
||||
deleted_rows = [local_map[identifier] for identifier in deleted_ids]
|
||||
await self.repository.delete_document_shadows(dataset_id, deleted_ids)
|
||||
await self.repository.update_stats(
|
||||
dataset_id,
|
||||
-len(deleted_ids),
|
||||
-sum(int(row.get("chunk_count") or 0) for row in deleted_rows),
|
||||
-sum(int(row.get("token_count") or 0) for row in deleted_rows),
|
||||
)
|
||||
await self.repository.session.commit()
|
||||
return new_count
|
||||
|
||||
async def sync_running(self) -> int:
|
||||
rows = await self.repository.running_documents()
|
||||
grouped: defaultdict[str, list[dict[str, Any]]] = defaultdict(list)
|
||||
for row in rows:
|
||||
grouped[str(row["dataset_id"])].append(row)
|
||||
updates = 0
|
||||
for dataset_id, documents in grouped.items():
|
||||
try:
|
||||
client = await self._client_for_dataset(dataset_id)
|
||||
except Exception:
|
||||
await self.repository.session.rollback()
|
||||
continue
|
||||
for local in documents:
|
||||
try:
|
||||
remote, _ = await client.documents(
|
||||
dataset_id, page=1, page_size=1, document_id=str(local["document_id"])
|
||||
)
|
||||
if not remote:
|
||||
await self.repository.mark_document_remote_deleted(
|
||||
str(local["document_id"]), shanghai_now_naive()
|
||||
)
|
||||
await self.repository.session.commit()
|
||||
updates += 1
|
||||
continue
|
||||
remote_status = remote[0].get("status")
|
||||
remote_run = remote[0].get("run")
|
||||
status_changed = remote_status is not None and str(remote_status) != str(local.get("status"))
|
||||
run_changed = remote_run is not None and str(remote_run) != str(local.get("run"))
|
||||
is_processing = remote_run in {"RUNNING", "UNSTART"}
|
||||
if not (status_changed or run_changed or is_processing):
|
||||
await self.repository.session.commit()
|
||||
continue
|
||||
before_tokens = int(local.get("token_count") or 0)
|
||||
await self.repository.sync_running_document(
|
||||
dataset_id,
|
||||
str(local["document_id"]),
|
||||
remote[0],
|
||||
shanghai_now_naive(),
|
||||
)
|
||||
delta = int(remote[0].get("token_count") or 0) - before_tokens
|
||||
if delta:
|
||||
await self.repository.update_stats(dataset_id, 0, 0, delta)
|
||||
await self.repository.session.commit()
|
||||
updates += 1
|
||||
except Exception:
|
||||
await self.repository.session.rollback()
|
||||
continue
|
||||
return updates
|
||||
|
||||
async def _client_for_dataset(self, dataset_id: str) -> RAGFlowClient:
|
||||
row = await self.repository.get_dataset(dataset_id)
|
||||
if row is None or not row.get("rag_model_id"):
|
||||
raise AppError(10164)
|
||||
return await self.datasets._client(str(row["rag_model_id"]))
|
||||
|
||||
|
||||
def _document_delete_error(exc: Exception, language: str | None) -> AppError:
|
||||
"""Match `new RenException(e.getMessage())` in the Java delete flow."""
|
||||
if isinstance(exc, AppError):
|
||||
message = exc.message or message_for(exc.code, language, *exc.params)
|
||||
else:
|
||||
message = str(exc)
|
||||
return AppError(500, message)
|
||||
|
||||
|
||||
async def _delete_cache_ignoring_errors(key: str) -> None:
|
||||
try:
|
||||
await get_redis().delete(key)
|
||||
except Exception:
|
||||
# The Java document cleanup and remote-missing cleanup explicitly log
|
||||
# and continue when Redis is unavailable.
|
||||
return
|
||||
@@ -0,0 +1,306 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import copy
|
||||
import json
|
||||
import uuid
|
||||
from typing import Any
|
||||
|
||||
from app.core.errors import AppError
|
||||
from app.core.redis import get_redis
|
||||
from app.core.security import AuthUser, shanghai_now_naive
|
||||
from app.repositories.model import ModelRepository, parse_json_object
|
||||
from app.schemas.model import ModelConfigBody, ModelProviderBody
|
||||
|
||||
SENSITIVE_FIELDS = {
|
||||
"api_key",
|
||||
"personal_access_token",
|
||||
"access_token",
|
||||
"token",
|
||||
"secret",
|
||||
"access_key_secret",
|
||||
"secret_key",
|
||||
}
|
||||
|
||||
|
||||
def _mask_middle(value: str) -> str:
|
||||
if not value.strip() or len(value) == 1:
|
||||
return value
|
||||
if len(value) <= 8:
|
||||
return value[:2] + "****" + value[-2:]
|
||||
return value[:4] + "*" * (len(value) - 8) + value[-4:]
|
||||
|
||||
|
||||
def mask_sensitive(value: Any) -> Any:
|
||||
if not isinstance(value, dict):
|
||||
return value
|
||||
result: dict[str, Any] = {}
|
||||
for key, item in value.items():
|
||||
if key.lower() in SENSITIVE_FIELDS and isinstance(item, str):
|
||||
result[key] = _mask_middle(item)
|
||||
elif isinstance(item, dict):
|
||||
result[key] = mask_sensitive(item)
|
||||
else:
|
||||
result[key] = copy.deepcopy(item)
|
||||
return result
|
||||
|
||||
|
||||
def _merge_config(original: dict[str, Any], updated: dict[str, Any]) -> dict[str, Any]:
|
||||
result = copy.deepcopy(original)
|
||||
for key, value in updated.items():
|
||||
if key.lower() in SENSITIVE_FIELDS:
|
||||
if isinstance(value, str) and "****" not in value:
|
||||
result[key] = value
|
||||
elif isinstance(value, dict):
|
||||
child = result.get(key)
|
||||
result[key] = _merge_config(child if isinstance(child, dict) else {}, value)
|
||||
else:
|
||||
result[key] = copy.deepcopy(value)
|
||||
for key in list(result):
|
||||
if key not in updated and key.lower() not in SENSITIVE_FIELDS:
|
||||
del result[key]
|
||||
return result
|
||||
|
||||
|
||||
def _model_dto(row: dict[str, Any], *, masked: bool = True) -> dict[str, Any]:
|
||||
config = parse_json_object(row.get("config_json"))
|
||||
return {
|
||||
"id": row.get("id"),
|
||||
"modelType": row.get("model_type"),
|
||||
"modelCode": row.get("model_code"),
|
||||
"modelName": row.get("model_name"),
|
||||
"isDefault": row.get("is_default"),
|
||||
"isEnabled": row.get("is_enabled"),
|
||||
"configJson": mask_sensitive(config) if masked else config,
|
||||
"docLink": row.get("doc_link"),
|
||||
"remark": row.get("remark"),
|
||||
"sort": row.get("sort"),
|
||||
}
|
||||
|
||||
|
||||
class ModelService:
|
||||
def __init__(self, repository: ModelRepository):
|
||||
self.repository = repository
|
||||
|
||||
async def names(self, model_type: str, model_name: str | None) -> list[dict[str, Any]]:
|
||||
return [
|
||||
{"id": row.get("id"), "modelName": row.get("model_name")}
|
||||
for row in await self.repository.list_model_names(model_type, model_name)
|
||||
]
|
||||
|
||||
async def llm_names(self, model_name: str | None) -> list[dict[str, Any]]:
|
||||
result: list[dict[str, Any]] = []
|
||||
for row in await self.repository.list_llm_names(model_name):
|
||||
config = parse_json_object(row.get("config_json")) or {}
|
||||
result.append(
|
||||
{"id": row.get("id"), "modelName": row.get("model_name"), "type": str(config.get("type", ""))}
|
||||
)
|
||||
return result
|
||||
|
||||
async def model_page(self, model_type: str, model_name: str | None, page: str, limit: str) -> dict[str, Any]:
|
||||
current, size = max(int(page), 1), int(limit)
|
||||
rows, total = await self.repository.list_model_configs(
|
||||
model_type=model_type,
|
||||
model_name=model_name,
|
||||
offset=(current - 1) * size,
|
||||
limit=size,
|
||||
)
|
||||
return {"total": total, "list": [_model_dto(row) for row in rows]}
|
||||
|
||||
async def get_model(self, model_id: str) -> dict[str, Any] | None:
|
||||
row = await self.repository.get_model(model_id)
|
||||
return _model_dto(row) if row else None
|
||||
|
||||
async def add(self, model_type: str, provider_code: str, body: ModelConfigBody) -> dict[str, Any]:
|
||||
if not model_type.strip() or not provider_code.strip():
|
||||
raise AppError(10131)
|
||||
model_id = body.id or uuid.uuid4().hex
|
||||
values = {
|
||||
"id": model_id,
|
||||
"model_type": model_type,
|
||||
"model_code": body.model_code,
|
||||
"model_name": body.model_name,
|
||||
"is_default": 0,
|
||||
"is_enabled": body.is_enabled,
|
||||
"config_json": json.dumps(body.config_json, ensure_ascii=False) if body.config_json is not None else None,
|
||||
"doc_link": body.doc_link,
|
||||
"remark": body.remark,
|
||||
"sort": body.sort,
|
||||
}
|
||||
async with self.repository.session.begin():
|
||||
# Keep the read and write in one transaction. A query before
|
||||
# ``begin()`` triggers SQLAlchemy autobegin and makes the explicit
|
||||
# transaction fail with InvalidRequestError.
|
||||
if await self.repository.get_provider(model_type, provider_code) is None:
|
||||
raise AppError(10162)
|
||||
await self.repository.insert_model(values)
|
||||
return _model_dto(values)
|
||||
|
||||
async def edit(
|
||||
self, model_type: str, provider_code: str, model_id: str, body: ModelConfigBody
|
||||
) -> dict[str, Any]:
|
||||
if not model_type.strip() or not provider_code.strip():
|
||||
raise AppError(10131)
|
||||
async with self.repository.session.begin():
|
||||
if await self.repository.get_provider(model_type, provider_code) is None:
|
||||
raise AppError(10162)
|
||||
original = await self.repository.get_model(model_id, for_update=True)
|
||||
if original is None:
|
||||
raise AppError(10051)
|
||||
updated_config = body.config_json
|
||||
if updated_config is not None and "llm" in updated_config:
|
||||
llm = await self.repository.get_model(str(updated_config["llm"]))
|
||||
llm_config = parse_json_object(llm.get("config_json")) if llm else None
|
||||
if llm is None or str(llm.get("model_type") or "").upper() != "LLM":
|
||||
raise AppError(10092)
|
||||
if llm_config and "type" in llm_config and llm_config["type"] not in {"openai", "ollama"}:
|
||||
raise AppError(10049)
|
||||
original_config = parse_json_object(original.get("config_json"))
|
||||
merged = (
|
||||
_merge_config(original_config, updated_config)
|
||||
if original_config is not None and updated_config is not None
|
||||
else original_config
|
||||
)
|
||||
values = {
|
||||
"id": model_id,
|
||||
"model_type": model_type,
|
||||
"model_code": original.get("model_code"),
|
||||
"model_name": body.model_name,
|
||||
"is_default": original.get("is_default"),
|
||||
"is_enabled": body.is_enabled,
|
||||
"config_json": json.dumps(merged, ensure_ascii=False) if merged is not None else None,
|
||||
"doc_link": original.get("doc_link"),
|
||||
"remark": body.remark,
|
||||
"sort": body.sort,
|
||||
}
|
||||
await self.repository.update_model(values)
|
||||
await self._clear_cache(model_id)
|
||||
return _model_dto(values)
|
||||
|
||||
async def delete(self, model_id: str) -> None:
|
||||
if not model_id.strip():
|
||||
raise AppError(10006)
|
||||
async with self.repository.session.begin():
|
||||
model = await self.repository.get_model(model_id, for_update=True)
|
||||
if model and int(model.get("is_default") or 0) == 1:
|
||||
raise AppError(10064)
|
||||
agents = await self.repository.model_agent_references(model_id)
|
||||
if agents:
|
||||
raise AppError(10093, params=("、".join(agents),))
|
||||
if model and str(model.get("model_type") or "").upper() == "LLM":
|
||||
if await self.repository.intent_reference_count(model_id):
|
||||
raise AppError(10094)
|
||||
await self.repository.delete_model(model_id)
|
||||
await self._clear_cache(model_id)
|
||||
|
||||
async def enable(self, model_id: str, status: int) -> str | None:
|
||||
async with self.repository.session.begin():
|
||||
model = await self.repository.get_model(model_id, for_update=True)
|
||||
if model is None:
|
||||
return "模型配置不存在"
|
||||
if status == 0 and int(model.get("is_default") or 0) > 0:
|
||||
return "默认模型配置不允许关闭"
|
||||
await self.repository.set_model_enabled(model_id, status)
|
||||
await self._clear_cache(model_id)
|
||||
return None
|
||||
|
||||
async def set_default(self, model_id: str) -> str | None:
|
||||
async with self.repository.session.begin():
|
||||
model = await self.repository.get_model(model_id, for_update=True)
|
||||
if model is None:
|
||||
return "模型配置不存在"
|
||||
model_type = str(model.get("model_type") or "")
|
||||
await self.repository.set_models_default(model_type, 0)
|
||||
await self.repository.execute(
|
||||
"UPDATE ai_model_config SET is_enabled=1, is_default=1 WHERE id=:id", {"id": model_id}
|
||||
)
|
||||
await self.repository.update_default_template_models(model_type, model_id)
|
||||
await self._clear_type_cache(model_type)
|
||||
return None
|
||||
|
||||
async def _clear_cache(self, model_id: str) -> None:
|
||||
redis = get_redis()
|
||||
await redis.delete(f"model:data:{model_id}", f"model:name:{model_id}")
|
||||
|
||||
async def _clear_type_cache(self, model_type: str) -> None:
|
||||
rows = await self.repository.fetch_all(
|
||||
"SELECT id FROM ai_model_config WHERE model_type=:type", {"type": model_type}
|
||||
)
|
||||
if rows:
|
||||
redis = get_redis()
|
||||
keys = [key for row in rows for key in (f"model:data:{row['id']}", f"model:name:{row['id']}")]
|
||||
await redis.delete(*keys)
|
||||
|
||||
|
||||
class ModelProviderService:
|
||||
def __init__(self, repository: ModelRepository):
|
||||
self.repository = repository
|
||||
|
||||
async def page(
|
||||
self, model_type: str | None, name: str | None, page: str, limit: str
|
||||
) -> dict[str, Any]:
|
||||
current, size = max(int(page), 1), int(limit)
|
||||
rows, total = await self.repository.list_providers(
|
||||
model_type=model_type, name=name, offset=(current - 1) * size, limit=size
|
||||
)
|
||||
return {"total": total, "list": rows}
|
||||
|
||||
@staticmethod
|
||||
def _validate(body: ModelProviderBody, *, update: bool) -> None:
|
||||
if update and (body.id is None or not body.id.strip()):
|
||||
raise AppError(10034, "id不能为空")
|
||||
for field, message in (
|
||||
(body.provider_code, "providerCode不能为空"),
|
||||
(body.model_type, "modelType不能为空"),
|
||||
(body.name, "name不能为空"),
|
||||
(body.fields, "fields(JSON格式)不能为空"),
|
||||
):
|
||||
if field is None or not field.strip():
|
||||
raise AppError(10034, message)
|
||||
if body.sort is None:
|
||||
raise AppError(10034, "sort不能为空")
|
||||
|
||||
async def add(self, body: ModelProviderBody, user: AuthUser) -> dict[str, Any]:
|
||||
self._validate(body, update=False)
|
||||
now = shanghai_now_naive()
|
||||
values = {
|
||||
"id": body.id or uuid.uuid4().hex,
|
||||
"model_type": body.model_type,
|
||||
"provider_code": body.provider_code,
|
||||
"name": body.name,
|
||||
"fields": body.fields,
|
||||
"sort": body.sort,
|
||||
"creator": user.id,
|
||||
"updater": user.id,
|
||||
"now": now,
|
||||
}
|
||||
async with self.repository.session.begin():
|
||||
await self.repository.insert_provider(values)
|
||||
return {
|
||||
# The Java service returns the request DTO, not the entity on which
|
||||
# MyBatis-Plus generated the UUID. Therefore an omitted id remains
|
||||
# null in the response even though the stored row has an id.
|
||||
"id": body.id, "modelType": body.model_type, "providerCode": body.provider_code,
|
||||
"name": body.name, "fields": body.fields, "sort": body.sort, "creator": user.id,
|
||||
"updater": user.id, "createDate": now, "updateDate": now,
|
||||
}
|
||||
|
||||
async def edit(self, body: ModelProviderBody, user: AuthUser) -> dict[str, Any]:
|
||||
self._validate(body, update=True)
|
||||
now = shanghai_now_naive()
|
||||
values = {
|
||||
"id": body.id, "model_type": body.model_type, "provider_code": body.provider_code,
|
||||
"name": body.name, "fields": body.fields, "sort": body.sort, "updater": user.id, "now": now,
|
||||
}
|
||||
async with self.repository.session.begin():
|
||||
if await self.repository.update_provider(values) == 0:
|
||||
raise AppError(10066)
|
||||
return {
|
||||
"id": body.id, "modelType": body.model_type, "providerCode": body.provider_code,
|
||||
"name": body.name, "fields": body.fields, "sort": body.sort, "updater": user.id,
|
||||
"updateDate": now, "creator": None, "createDate": None,
|
||||
}
|
||||
|
||||
async def delete(self, ids: list[str]) -> None:
|
||||
async with self.repository.session.begin():
|
||||
if await self.repository.delete_providers(ids) == 0:
|
||||
raise AppError(10043)
|
||||
@@ -0,0 +1,486 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import base64
|
||||
import hashlib
|
||||
import hmac
|
||||
import io
|
||||
import json
|
||||
import logging
|
||||
import re
|
||||
import secrets
|
||||
import string
|
||||
import time
|
||||
import urllib.parse
|
||||
import uuid
|
||||
from datetime import datetime, timedelta
|
||||
from typing import Any, Protocol, cast
|
||||
|
||||
import httpx
|
||||
from fastapi import Request
|
||||
from PIL import Image, ImageDraw, ImageFont
|
||||
from redis.asyncio import Redis
|
||||
|
||||
from app.core.config import get_settings
|
||||
from app.core.crypto import bcrypt_hash, bcrypt_matches, generate_database_token, sm2_decrypt_c1c3c2
|
||||
from app.core.errors import AppError, ErrorCode
|
||||
from app.core.ids import snowflake
|
||||
from app.core.redis import JavaRedisCodec, get_redis
|
||||
from app.core.security import AuthUser, shanghai_now_naive
|
||||
from app.repositories.security import SecurityRepository
|
||||
from app.schemas.security import (
|
||||
LoginRequest,
|
||||
PasswordChangeRequest,
|
||||
RetrievePasswordRequest,
|
||||
SmsVerificationRequest,
|
||||
)
|
||||
from app.services.java_validation import validation_message
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
TOKEN_EXPIRE_SECONDS = 12 * 60 * 60
|
||||
CAPTCHA_TTL_SECONDS = 5 * 60
|
||||
CAPTCHA_LENGTH = 5
|
||||
PHONE_PATTERN = re.compile(r"^\+[1-9]\d{0,3}[1-9]\d{4,14}$")
|
||||
STRONG_PASSWORD = re.compile(r"^(?=.*[0-9])(?=.*[a-z])(?=.*[A-Z]).+$")
|
||||
|
||||
|
||||
class SmsSender(Protocol):
|
||||
async def send_verification_code(self, phone: str | None, code: str) -> None: ...
|
||||
|
||||
|
||||
class AliyunSmsSender:
|
||||
"""Minimal implementation of the Aliyun Dysmsapi RPC request used by the Java SDK."""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
repository: SecurityRepository,
|
||||
*,
|
||||
redis: Redis | None = None,
|
||||
client: httpx.AsyncClient | None = None,
|
||||
endpoint: str = "https://dysmsapi.aliyuncs.com/",
|
||||
):
|
||||
self.repository = repository
|
||||
self.redis = redis or get_redis()
|
||||
self.client = client
|
||||
self.endpoint = endpoint
|
||||
|
||||
async def send_verification_code(self, phone: str | None, code: str) -> None:
|
||||
access_key_id = await self._param("aliyun.sms.access_key_id") or ""
|
||||
access_key_secret = await self._param("aliyun.sms.access_key_secret") or ""
|
||||
sign_name = await self._param("aliyun.sms.sign_name") or ""
|
||||
template_code = await self._param("aliyun.sms.sms_code_template_code") or ""
|
||||
# The Tea SDK constructs its client before the refundable send block;
|
||||
# blank credentials therefore map to SMS_CONNECTION_FAILED (10056).
|
||||
if not access_key_id.strip() or not access_key_secret.strip():
|
||||
raise AppError(10056)
|
||||
timestamp = time.strftime("%Y-%m-%dT%H:%M:%SZ", time.gmtime())
|
||||
params: dict[str, str] = {
|
||||
"AccessKeyId": access_key_id,
|
||||
"Action": "SendSms",
|
||||
"Format": "JSON",
|
||||
"RegionId": "cn-hangzhou",
|
||||
"SignatureMethod": "HMAC-SHA1",
|
||||
"SignatureNonce": str(uuid.uuid4()),
|
||||
"SignatureVersion": "1.0",
|
||||
"SignName": sign_name,
|
||||
"TemplateCode": template_code,
|
||||
"TemplateParam": json.dumps({"code": code}, ensure_ascii=False, separators=(",", ":")),
|
||||
"Timestamp": timestamp,
|
||||
"Version": "2017-05-25",
|
||||
}
|
||||
if phone is not None:
|
||||
params["PhoneNumbers"] = phone
|
||||
params["Signature"] = self._signature(params, access_key_secret)
|
||||
if self.client is not None:
|
||||
response = await self.client.post(self.endpoint, data=params)
|
||||
response.raise_for_status()
|
||||
return
|
||||
timeout = get_settings().external_request_timeout_seconds
|
||||
async with httpx.AsyncClient(timeout=timeout) as client:
|
||||
response = await client.post(self.endpoint, data=params)
|
||||
response.raise_for_status()
|
||||
|
||||
async def _param(self, code: str) -> str | None:
|
||||
cached = JavaRedisCodec.decode(await cast(Any, self.redis.hget)("sys:params", code))
|
||||
if cached is not None:
|
||||
return str(cached)
|
||||
value = await self.repository.get_param_value(code)
|
||||
if value is not None:
|
||||
await cast(Any, self.redis.hset)("sys:params", code, JavaRedisCodec.encode(value))
|
||||
await cast(Any, self.redis.expire)("sys:params", 24 * 60 * 60)
|
||||
return value
|
||||
|
||||
@classmethod
|
||||
def _signature(cls, params: dict[str, str], secret: str) -> str:
|
||||
canonical = "&".join(
|
||||
f"{cls._percent_encode(key)}={cls._percent_encode(value)}" for key, value in sorted(params.items())
|
||||
)
|
||||
string_to_sign = f"POST&%2F&{cls._percent_encode(canonical)}"
|
||||
digest = hmac.new(
|
||||
f"{secret}&".encode(),
|
||||
string_to_sign.encode(),
|
||||
digestmod=hashlib.sha1, # noqa: S324 - mandated by Aliyun RPC SignatureMethod
|
||||
).digest()
|
||||
return base64.b64encode(digest).decode("ascii")
|
||||
|
||||
@staticmethod
|
||||
def _percent_encode(value: str) -> str:
|
||||
return urllib.parse.quote(str(value), safe="~")
|
||||
|
||||
|
||||
class CaptchaService:
|
||||
def __init__(self, redis: Redis | None = None):
|
||||
self.redis = redis or get_redis()
|
||||
|
||||
async def create(self, identifier: str) -> bytes:
|
||||
code = "".join(secrets.choice(string.ascii_letters + string.digits) for _ in range(CAPTCHA_LENGTH))
|
||||
await self._set_cache(identifier, code)
|
||||
return self._render_gif(code)
|
||||
|
||||
async def validate(self, identifier: str | None, code: str | None, *, delete: bool) -> bool:
|
||||
if not code or not code.strip():
|
||||
return False
|
||||
key = self._captcha_key(identifier)
|
||||
cached = JavaRedisCodec.decode(await cast(Any, self.redis.get(key)))
|
||||
if cached is not None and delete:
|
||||
await cast(Any, self.redis.delete(key))
|
||||
return cached is not None and code.casefold() == str(cached).casefold()
|
||||
|
||||
async def set_sms_code(self, phone: str | None, code: str) -> None:
|
||||
await self._set_cache(f"sms:Validate:Code:{phone}", code)
|
||||
|
||||
async def validate_sms_code(self, phone: str | None, code: str | None, *, delete: bool = False) -> bool:
|
||||
return await self.validate(f"sms:Validate:Code:{phone}", code, delete=delete)
|
||||
|
||||
async def _set_cache(self, identifier: str, value: str) -> None:
|
||||
await cast(Any, self.redis.set)(
|
||||
self._captcha_key(identifier),
|
||||
JavaRedisCodec.encode(value),
|
||||
ex=CAPTCHA_TTL_SECONDS,
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _captcha_key(identifier: str | None) -> str:
|
||||
return f"sys:captcha:{'null' if identifier is None else identifier}"
|
||||
|
||||
@staticmethod
|
||||
def _render_gif(code: str) -> bytes:
|
||||
image = Image.new("RGB", (150, 40), (248, 248, 248))
|
||||
draw = ImageDraw.Draw(image)
|
||||
for _ in range(8):
|
||||
color = tuple(secrets.randbelow(150) for _ in range(3))
|
||||
draw.line(
|
||||
(
|
||||
secrets.randbelow(150),
|
||||
secrets.randbelow(40),
|
||||
secrets.randbelow(150),
|
||||
secrets.randbelow(40),
|
||||
),
|
||||
fill=color,
|
||||
width=1,
|
||||
)
|
||||
font = ImageFont.load_default(size=24)
|
||||
for index, character in enumerate(code):
|
||||
color = tuple(secrets.randbelow(120) for _ in range(3))
|
||||
draw.text((10 + index * 27, 6 + secrets.randbelow(5)), character, font=font, fill=color)
|
||||
output = io.BytesIO()
|
||||
image.save(output, format="GIF")
|
||||
return output.getvalue()
|
||||
|
||||
|
||||
class SecurityService:
|
||||
def __init__(
|
||||
self,
|
||||
repository: SecurityRepository,
|
||||
*,
|
||||
redis: Redis | None = None,
|
||||
captcha: CaptchaService | None = None,
|
||||
sms_sender: SmsSender | None = None,
|
||||
):
|
||||
self.repository = repository
|
||||
self.redis = redis or get_redis()
|
||||
self.captcha = captcha or CaptchaService(self.redis)
|
||||
self.sms_sender = sms_sender or AliyunSmsSender(repository, redis=self.redis)
|
||||
|
||||
async def login(self, dto: LoginRequest, request: Request) -> dict[str, Any]:
|
||||
password = await self._decrypt_and_validate_captcha(dto.password, dto.captcha_id)
|
||||
user = await self.repository.get_user_by_username(dto.username)
|
||||
if user is None or not bcrypt_matches(password, cast(str | None, user.get("password"))):
|
||||
raise AppError(ErrorCode.ACCOUNT_PASSWORD_ERROR)
|
||||
token = await self._create_token(int(user["id"]))
|
||||
await self.repository.session.commit()
|
||||
return {
|
||||
"token": token,
|
||||
"expire": TOKEN_EXPIRE_SECONDS,
|
||||
"clientHash": self._client_hash(request),
|
||||
}
|
||||
|
||||
async def register(self, dto: LoginRequest) -> None:
|
||||
if not await self.allow_user_register():
|
||||
raise AppError(10072)
|
||||
password = await self._decrypt_and_validate_captcha(dto.password, dto.captcha_id)
|
||||
if await self._mobile_registration_enabled():
|
||||
if dto.username is None or not PHONE_PATTERN.fullmatch(dto.username):
|
||||
raise AppError(10069)
|
||||
if not await self.captcha.validate_sms_code(dto.username, dto.mobile_captcha, delete=False):
|
||||
raise AppError(10075)
|
||||
if await self.repository.get_user_by_username(dto.username) is not None:
|
||||
raise AppError(10070)
|
||||
if not STRONG_PASSWORD.fullmatch(password):
|
||||
raise AppError(ErrorCode.PASSWORD_WEAK_ERROR)
|
||||
now = shanghai_now_naive()
|
||||
user_count = await self.repository.count_users()
|
||||
await self.repository.insert_user(
|
||||
user_id=snowflake.next_id(),
|
||||
username=dto.username,
|
||||
password=bcrypt_hash(password),
|
||||
super_admin=1 if user_count == 0 else 0,
|
||||
now=now,
|
||||
)
|
||||
await self.repository.session.commit()
|
||||
|
||||
async def change_password(
|
||||
self,
|
||||
user: AuthUser,
|
||||
dto: PasswordChangeRequest,
|
||||
accept_language: str | None = None,
|
||||
) -> None:
|
||||
self._require_not_blank(dto.password, "sysuser.password.require", accept_language)
|
||||
self._require_not_blank(dto.new_password, "sysuser.password.require", accept_language)
|
||||
assert dto.password is not None
|
||||
assert dto.new_password is not None
|
||||
row = await self.repository.get_user_by_id(user.id)
|
||||
if row is None:
|
||||
raise AppError(ErrorCode.TOKEN_INVALID)
|
||||
if not bcrypt_matches(dto.password, cast(str | None, row.get("password"))):
|
||||
raise AppError(10048)
|
||||
if not STRONG_PASSWORD.fullmatch(dto.new_password):
|
||||
raise AppError(ErrorCode.PASSWORD_WEAK_ERROR)
|
||||
now = shanghai_now_naive()
|
||||
await self.repository.update_password(
|
||||
user.id,
|
||||
bcrypt_hash(dto.new_password),
|
||||
now,
|
||||
preserve_audit_fields=True,
|
||||
)
|
||||
# SysUserService.changePassword commits before the non-transactional token service logs out.
|
||||
await self.repository.session.commit()
|
||||
await self.repository.expire_user_token(user.id, now - timedelta(minutes=1))
|
||||
await self.repository.session.commit()
|
||||
|
||||
async def retrieve_password(
|
||||
self,
|
||||
dto: RetrievePasswordRequest,
|
||||
accept_language: str | None = None,
|
||||
) -> None:
|
||||
if not await self._mobile_registration_enabled():
|
||||
raise AppError(10073)
|
||||
self._require_not_blank(dto.phone, "sysuser.password.require", accept_language)
|
||||
self._require_not_blank(dto.code, "sysuser.password.require", accept_language)
|
||||
self._require_not_blank(dto.password, "sysuser.password.require", accept_language)
|
||||
self._require_not_blank(dto.captcha_id, "sysuser.uuid.require", accept_language)
|
||||
assert dto.phone is not None
|
||||
assert dto.code is not None
|
||||
assert dto.password is not None
|
||||
assert dto.captcha_id is not None
|
||||
if not PHONE_PATTERN.fullmatch(dto.phone):
|
||||
raise AppError(10074)
|
||||
user = await self.repository.get_user_by_username(dto.phone)
|
||||
if user is None:
|
||||
raise AppError(10071)
|
||||
if not await self.captcha.validate_sms_code(dto.phone, dto.code, delete=False):
|
||||
raise AppError(10075)
|
||||
password = await self._decrypt_and_validate_captcha(dto.password, dto.captcha_id)
|
||||
if not STRONG_PASSWORD.fullmatch(password):
|
||||
raise AppError(ErrorCode.PASSWORD_WEAK_ERROR)
|
||||
await self.repository.update_password(int(user["id"]), bcrypt_hash(password), shanghai_now_naive())
|
||||
await self.repository.session.commit()
|
||||
|
||||
async def send_sms_verification(self, dto: SmsVerificationRequest) -> None:
|
||||
if not await self.captcha.validate(dto.captcha_id, dto.captcha, delete=False):
|
||||
raise AppError(10067)
|
||||
if not await self._mobile_registration_enabled():
|
||||
raise AppError(10068)
|
||||
phone_key = "null" if dto.phone is None else dto.phone
|
||||
last_send_key = f"sms:Validate:Code:{phone_key}:last_send_time"
|
||||
current_ms = int(time.time() * 1000)
|
||||
created = await cast(Any, self.redis.set)(last_send_key, str(current_ms), ex=60, nx=True)
|
||||
if not created:
|
||||
raw_last = await cast(Any, self.redis.get)(last_send_key)
|
||||
if raw_last is not None:
|
||||
last_ms = int(raw_last.decode() if isinstance(raw_last, bytes) else raw_last)
|
||||
difference = current_ms - last_ms
|
||||
if difference < 60_000:
|
||||
raise AppError(10060, params=(str(max(0, (60_000 - difference) // 1000)),))
|
||||
|
||||
today_key = f"sms:Validate:Code:{phone_key}:today_count"
|
||||
raw_count = await cast(Any, self.redis.get)(today_key)
|
||||
decoded_count = JavaRedisCodec.decode(raw_count)
|
||||
today_count = int(decoded_count or 0)
|
||||
raw_maximum = await self._get_param("server.sms_max_send_count", from_cache=True)
|
||||
maximum = int(raw_maximum) if raw_maximum is not None and raw_maximum != "" else 5
|
||||
if today_count >= maximum:
|
||||
raise AppError(10047)
|
||||
|
||||
code = "".join(secrets.choice(string.digits) for _ in range(6))
|
||||
await self.captcha.set_sms_code(dto.phone, code)
|
||||
new_count = await cast(Any, self.redis.incr)(today_key)
|
||||
if int(new_count) == 1:
|
||||
await cast(Any, self.redis.expire)(today_key, 24 * 60 * 60)
|
||||
try:
|
||||
await self.sms_sender.send_verification_code(dto.phone, code)
|
||||
except AppError:
|
||||
# Java raises connection-construction failures before entering its refundable send attempt.
|
||||
raise
|
||||
except Exception as exc:
|
||||
logger.warning("Aliyun SMS request failed", exc_info=exc)
|
||||
await cast(Any, self.redis.delete)(today_key)
|
||||
raise AppError(10055) from exc
|
||||
|
||||
async def public_config(self) -> dict[str, Any]:
|
||||
public_key = await self._get_param("server.public_key", from_cache=True)
|
||||
if public_key is None or not public_key.strip():
|
||||
raise AppError(10129)
|
||||
menu_config = await self._get_param("system-web.menu", from_cache=True)
|
||||
result: dict[str, Any] = {
|
||||
"enableMobileRegister": await self._mobile_registration_enabled(),
|
||||
"version": "0.9.5",
|
||||
"year": f"©{shanghai_now_naive().year}",
|
||||
"allowUserRegister": await self.allow_user_register(),
|
||||
"mobileAreaList": await self._dict_data_by_type("MOBILE_AREA"),
|
||||
"beianIcpNum": await self._get_param("server.beian_icp_num", from_cache=True),
|
||||
"beianGaNum": await self._get_param("server.beian_ga_num", from_cache=True),
|
||||
"name": await self._get_param("server.name", from_cache=True),
|
||||
"sm2PublicKey": public_key,
|
||||
}
|
||||
if menu_config is not None and menu_config.strip():
|
||||
result["systemWebMenu"] = json.loads(menu_config)
|
||||
return result
|
||||
|
||||
async def allow_user_register(self) -> bool:
|
||||
value = await self._get_param("server.allow_user_register", from_cache=True)
|
||||
if value == "true":
|
||||
return True
|
||||
return await self.repository.count_users() == 0
|
||||
|
||||
async def _create_token(self, user_id: int) -> str:
|
||||
now = shanghai_now_naive()
|
||||
expire_date = now + timedelta(seconds=TOKEN_EXPIRE_SECONDS)
|
||||
current = await self.repository.get_token_by_user_id(user_id, for_update=True)
|
||||
if current is None:
|
||||
token = generate_database_token()
|
||||
await self.repository.insert_token(
|
||||
token_id=snowflake.next_id(),
|
||||
user_id=user_id,
|
||||
token=token,
|
||||
now=now,
|
||||
expire_date=expire_date,
|
||||
)
|
||||
return token
|
||||
stored_expiry = self._datetime(current.get("expire_date"))
|
||||
token = str(current["token"])
|
||||
if stored_expiry is None or stored_expiry < now:
|
||||
token = generate_database_token()
|
||||
await self.repository.update_token(
|
||||
token_id=int(current["id"]),
|
||||
token=token,
|
||||
now=now,
|
||||
expire_date=expire_date,
|
||||
)
|
||||
return token
|
||||
|
||||
async def _decrypt_and_validate_captcha(
|
||||
self,
|
||||
encrypted_password: str | None,
|
||||
captcha_id: str | None,
|
||||
) -> str:
|
||||
private_key = await self._get_param("server.private_key", from_cache=True)
|
||||
if private_key is None or not private_key.strip():
|
||||
raise AppError(10129)
|
||||
try:
|
||||
if encrypted_password is None:
|
||||
raise ValueError("encrypted password is null")
|
||||
content = sm2_decrypt_c1c3c2(private_key, encrypted_password)
|
||||
except Exception as exc:
|
||||
raise AppError(10130) from exc
|
||||
if len(content) > CAPTCHA_LENGTH:
|
||||
embedded_captcha = content[:CAPTCHA_LENGTH]
|
||||
if not await self.captcha.validate(captcha_id, embedded_captcha, delete=True):
|
||||
raise AppError(10067)
|
||||
return content[CAPTCHA_LENGTH:]
|
||||
if content:
|
||||
raise AppError(10067)
|
||||
raise AppError(10130)
|
||||
|
||||
async def _mobile_registration_enabled(self) -> bool:
|
||||
value = await self._get_param("server.enable_mobile_register", from_cache=True)
|
||||
if value is None or not value.strip():
|
||||
return False
|
||||
try:
|
||||
parsed = json.loads(value.lower())
|
||||
except json.JSONDecodeError as exc:
|
||||
raise AppError(ErrorCode.PARAMS_GET_ERROR) from exc
|
||||
return bool(parsed)
|
||||
|
||||
async def _get_param(self, code: str, *, from_cache: bool) -> str | None:
|
||||
if from_cache:
|
||||
cached = JavaRedisCodec.decode(await cast(Any, self.redis.hget)("sys:params", code))
|
||||
if cached is not None:
|
||||
return str(cached)
|
||||
value = await self.repository.get_param_value(code)
|
||||
if from_cache and value is not None:
|
||||
await cast(Any, self.redis.hset)("sys:params", code, JavaRedisCodec.encode(value))
|
||||
await cast(Any, self.redis.expire)("sys:params", 24 * 60 * 60)
|
||||
return value
|
||||
|
||||
async def _dict_data_by_type(self, dict_type: str) -> list[dict[str, Any]]:
|
||||
key = f"sys:dict:data:{dict_type}"
|
||||
cached = JavaRedisCodec.decode(await cast(Any, self.redis.get)(key))
|
||||
if isinstance(cached, list):
|
||||
return cast(list[dict[str, Any]], cached)
|
||||
values = await self.repository.get_mobile_area_items()
|
||||
await cast(Any, self.redis.set)(
|
||||
key,
|
||||
JavaRedisCodec.encode(
|
||||
values,
|
||||
item_java_type="xiaozhi.modules.sys.vo.SysDictDataItem",
|
||||
),
|
||||
ex=24 * 60 * 60,
|
||||
)
|
||||
return values
|
||||
|
||||
@staticmethod
|
||||
def _client_hash(request: Request) -> str:
|
||||
user_agent = request.headers.get("User-Agent", "").lower()
|
||||
forwarded_headers = (
|
||||
"x-forwarded-for",
|
||||
"Proxy-Client-IP",
|
||||
"WL-Proxy-Client-IP",
|
||||
"HTTP_CLIENT_IP",
|
||||
"HTTP_X_FORWARDED_FOR",
|
||||
)
|
||||
ip_address = next(
|
||||
(
|
||||
value
|
||||
for header in forwarded_headers
|
||||
if (value := request.headers.get(header)) and value.casefold() != "unknown"
|
||||
),
|
||||
request.client.host if request.client else "",
|
||||
)
|
||||
date = shanghai_now_naive().strftime("%Y-%m-%d")
|
||||
return hashlib.md5( # noqa: S324 - Java clientHash compatibility requires MD5
|
||||
f"{ip_address}{date}{user_agent}".encode(), usedforsecurity=False
|
||||
).hexdigest()
|
||||
|
||||
@staticmethod
|
||||
def _datetime(value: Any) -> datetime | None:
|
||||
if value is None or isinstance(value, datetime):
|
||||
return value
|
||||
if isinstance(value, str):
|
||||
return datetime.fromisoformat(value)
|
||||
raise TypeError(f"Unsupported database datetime value: {type(value).__name__}")
|
||||
|
||||
@staticmethod
|
||||
def _require_not_blank(value: str | None, key: str, accept_language: str | None) -> None:
|
||||
if value is None or not value.strip():
|
||||
raise AppError(500, validation_message(key, accept_language))
|
||||
@@ -0,0 +1,715 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import base64
|
||||
import hashlib
|
||||
import hmac
|
||||
import json
|
||||
import logging
|
||||
import re
|
||||
import secrets
|
||||
import string
|
||||
import time
|
||||
import uuid
|
||||
from datetime import datetime
|
||||
from typing import Any, cast
|
||||
from zoneinfo import ZoneInfo
|
||||
|
||||
import httpx
|
||||
from redis.asyncio import Redis
|
||||
from websockets.asyncio.client import connect
|
||||
|
||||
from app.core.config import get_settings
|
||||
from app.core.crypto import bcrypt_hash
|
||||
from app.core.errors import AppError, ErrorCode
|
||||
from app.core.ids import snowflake
|
||||
from app.core.redis import JavaRedisCodec, get_redis
|
||||
from app.core.security import AuthUser, shanghai_now_naive
|
||||
from app.repositories.sys import SysRepository
|
||||
from app.schemas.sys import DictDataPayload, DictTypePayload, EmitServerActionRequest, SysParamPayload
|
||||
from app.services.java_validation import validation_message
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
WS_PATTERN = re.compile(r"^wss?://[\w.-]+(?:\.[\w.-]+)*(?::\d+)?(?:/[\w.-]*)*$")
|
||||
|
||||
|
||||
class AdminService:
|
||||
def __init__(self, repository: SysRepository):
|
||||
self.repository = repository
|
||||
|
||||
async def page_users(self, *, mobile: str | None, page: int, limit: int) -> dict[str, Any]:
|
||||
rows, total = await self.repository.page_users(
|
||||
mobile=mobile,
|
||||
page=max(1, page),
|
||||
limit=max(0, limit),
|
||||
)
|
||||
values = [
|
||||
{
|
||||
"deviceCount": str(row.get("device_count") or 0),
|
||||
"mobile": row.get("username"),
|
||||
"status": row.get("status"),
|
||||
"userid": str(row["id"]),
|
||||
"createDate": row.get("create_date"),
|
||||
}
|
||||
for row in rows
|
||||
]
|
||||
return {"list": values, "total": total}
|
||||
|
||||
async def reset_password(self, user_id: int, user: AuthUser) -> str:
|
||||
password = self._generate_password()
|
||||
await self.repository.reset_user_password(user_id, bcrypt_hash(password), user.id, shanghai_now_naive())
|
||||
await self.repository.session.commit()
|
||||
return password
|
||||
|
||||
async def delete_user(self, user_id: int) -> None:
|
||||
try:
|
||||
await self.repository.delete_user_cascade(user_id)
|
||||
await self.repository.session.commit()
|
||||
except Exception:
|
||||
await self.repository.session.rollback()
|
||||
raise
|
||||
|
||||
async def change_status(self, status: int, user_ids: list[str], user: AuthUser) -> None:
|
||||
# SysUserServiceImpl.changeStatus has an outer Spring transaction: a later
|
||||
# parse/update failure rolls back every earlier item in the same request.
|
||||
try:
|
||||
for value in user_ids:
|
||||
await self.repository.change_user_status(status, [int(value)], user.id, shanghai_now_naive())
|
||||
await self.repository.session.commit()
|
||||
except Exception:
|
||||
await self.repository.session.rollback()
|
||||
raise
|
||||
|
||||
async def page_devices(self, *, keywords: str | None, page: int, limit: int) -> dict[str, Any]:
|
||||
rows, total = await self.repository.page_devices(
|
||||
keywords=keywords,
|
||||
page=max(1, page),
|
||||
limit=max(0, limit),
|
||||
)
|
||||
result = []
|
||||
for row in rows:
|
||||
result.append(
|
||||
{
|
||||
"appVersion": row.get("app_version"),
|
||||
"bindUserName": row.get("bind_user_name"),
|
||||
"deviceType": row.get("board"),
|
||||
"board": row.get("board"),
|
||||
"id": row.get("id"),
|
||||
"macAddress": row.get("mac_address"),
|
||||
"alias": row.get("alias"),
|
||||
"otaUpgrade": None,
|
||||
"recentChatTime": self._short_time(cast(datetime | str | None, row.get("update_date"))),
|
||||
"lastConnectedAtTimestamp": self._timestamp_ms(
|
||||
cast(datetime | str | None, row.get("last_connected_at"))
|
||||
),
|
||||
"createDateTimestamp": self._timestamp_ms(
|
||||
cast(datetime | str | None, row.get("create_date"))
|
||||
),
|
||||
"createDate": self._utc_datetime_string(
|
||||
cast(datetime | str | None, row.get("create_date"))
|
||||
),
|
||||
}
|
||||
)
|
||||
return {"list": result, "total": total}
|
||||
|
||||
@staticmethod
|
||||
def _generate_password() -> str:
|
||||
characters = string.ascii_letters + string.digits + "!@#$%^&*()"
|
||||
values = [
|
||||
secrets.choice(string.digits),
|
||||
secrets.choice(string.ascii_lowercase),
|
||||
secrets.choice(string.ascii_uppercase),
|
||||
secrets.choice("!@#$%^&*()"),
|
||||
]
|
||||
values.extend(secrets.choice(characters) for _ in range(8))
|
||||
secrets.SystemRandom().shuffle(values)
|
||||
return "".join(values)
|
||||
|
||||
@staticmethod
|
||||
def _timestamp_ms(value: datetime | str | None) -> int | None:
|
||||
value = AdminService._database_datetime(value)
|
||||
if value is None:
|
||||
return None
|
||||
timezone = ZoneInfo(get_settings().timezone)
|
||||
localized = value if value.tzinfo else value.replace(tzinfo=timezone)
|
||||
return int(localized.timestamp() * 1000)
|
||||
|
||||
@staticmethod
|
||||
def _short_time(value: datetime | str | None) -> str | None:
|
||||
value = AdminService._database_datetime(value)
|
||||
if value is None:
|
||||
return None
|
||||
now = shanghai_now_naive()
|
||||
if value.tzinfo:
|
||||
value = value.astimezone(ZoneInfo(get_settings().timezone)).replace(tzinfo=None)
|
||||
seconds = int((now - value).total_seconds())
|
||||
if seconds <= 10:
|
||||
return "刚刚"
|
||||
if seconds < 60:
|
||||
return f"{seconds}秒前"
|
||||
if seconds < 3600:
|
||||
return f"{seconds // 60}分钟前"
|
||||
if seconds < 86400:
|
||||
return f"{seconds // 3600}小时前"
|
||||
if seconds < 604800:
|
||||
return f"{seconds // 86400}天前"
|
||||
return value.strftime("%Y-%m-%d %H:%M:%S")
|
||||
|
||||
@staticmethod
|
||||
def _utc_datetime_string(value: datetime | str | None) -> str | None:
|
||||
value = AdminService._database_datetime(value)
|
||||
if value is None:
|
||||
return None
|
||||
timezone = ZoneInfo(get_settings().timezone)
|
||||
localized = value if value.tzinfo else value.replace(tzinfo=timezone)
|
||||
return localized.astimezone(ZoneInfo("UTC")).strftime("%Y-%m-%d %H:%M:%S")
|
||||
|
||||
@staticmethod
|
||||
def _database_datetime(value: datetime | str | None) -> datetime | None:
|
||||
if isinstance(value, str):
|
||||
return datetime.fromisoformat(value)
|
||||
return value
|
||||
|
||||
|
||||
class ParamExternalValidator:
|
||||
def __init__(self, client: httpx.AsyncClient | None = None):
|
||||
self.client = client
|
||||
|
||||
async def validate(self, code: str, value: str) -> None:
|
||||
if code == "server.websocket":
|
||||
await self._websockets(value)
|
||||
elif code == "server.ota":
|
||||
await self._http_endpoint(value, kind="ota")
|
||||
elif code == "server.mcp_endpoint":
|
||||
await self._http_endpoint(value, kind="mcp")
|
||||
elif code == "server.voice_print":
|
||||
await self._http_endpoint(value, kind="voiceprint")
|
||||
elif code == "server.mqtt_signature_key":
|
||||
self._mqtt_secret(value)
|
||||
|
||||
async def _websockets(self, value: str) -> None:
|
||||
urls = value.split(";")
|
||||
while urls and urls[-1] == "":
|
||||
urls.pop()
|
||||
if not urls:
|
||||
raise AppError(10098)
|
||||
for raw_url in urls:
|
||||
if not raw_url.strip():
|
||||
continue
|
||||
if "localhost" in raw_url or "127.0.0.1" in raw_url:
|
||||
raise AppError(10099)
|
||||
if not WS_PATTERN.fullmatch(raw_url.strip()):
|
||||
raise AppError(10100)
|
||||
try:
|
||||
async with connect(raw_url, open_timeout=5):
|
||||
pass
|
||||
except Exception as exc:
|
||||
raise AppError(10101) from exc
|
||||
|
||||
async def _http_endpoint(self, value: str, *, kind: str) -> None:
|
||||
if not value.strip() or value == "null":
|
||||
return
|
||||
if "localhost" in value or "127.0.0.1" in value:
|
||||
raise AppError({"ota": 10103, "mcp": 10110, "voiceprint": 10116}[kind])
|
||||
if kind == "ota":
|
||||
if not value.lower().startswith("http"):
|
||||
raise AppError(10104)
|
||||
if not value.endswith("/ota/"):
|
||||
raise AppError(10105)
|
||||
elif kind == "mcp":
|
||||
if "key" not in value.lower():
|
||||
raise AppError(10111)
|
||||
else:
|
||||
if "key" not in value.lower():
|
||||
raise AppError(10117)
|
||||
if not value.lower().startswith("http"):
|
||||
raise AppError(10118)
|
||||
|
||||
final_code = {"ota": 10108, "mcp": 10114, "voiceprint": 10121}[kind]
|
||||
marker = {"ota": "OTA", "mcp": "success", "voiceprint": "healthy"}[kind]
|
||||
try:
|
||||
if self.client is not None:
|
||||
response = await self.client.get(value)
|
||||
else:
|
||||
async with httpx.AsyncClient(timeout=get_settings().external_request_timeout_seconds) as client:
|
||||
response = await client.get(value)
|
||||
if response.status_code != 200 or marker not in response.text:
|
||||
raise ValueError("external endpoint response did not match Java validation")
|
||||
except Exception as exc:
|
||||
raise AppError(final_code) from exc
|
||||
|
||||
@staticmethod
|
||||
def _mqtt_secret(secret: str) -> None:
|
||||
if not secret.strip() or secret == "null": # noqa: S105 - sentinel value from the Java parameter table
|
||||
raise AppError(10122)
|
||||
if len(secret) < 8:
|
||||
raise AppError(10123)
|
||||
if not re.search(r"[a-z]", secret) or not re.search(r"[A-Z]", secret):
|
||||
raise AppError(10124)
|
||||
lowered = secret.lower()
|
||||
if any(weak in lowered for weak in ("test", "1234", "admin", "password", "qwerty", "xiaozhi")):
|
||||
raise AppError(10125)
|
||||
|
||||
|
||||
class SysParamService:
|
||||
def __init__(
|
||||
self,
|
||||
repository: SysRepository,
|
||||
*,
|
||||
redis: Redis | None = None,
|
||||
validator: ParamExternalValidator | None = None,
|
||||
):
|
||||
self.repository = repository
|
||||
self.redis = redis or get_redis()
|
||||
self.validator = validator or ParamExternalValidator()
|
||||
|
||||
async def page(
|
||||
self,
|
||||
*,
|
||||
param_code: str | None,
|
||||
page: int,
|
||||
limit: int,
|
||||
order_field: str | None,
|
||||
order: str | None,
|
||||
) -> dict[str, Any]:
|
||||
rows, total = await self.repository.page_params(
|
||||
param_code=param_code,
|
||||
page=max(1, page),
|
||||
limit=max(0, limit),
|
||||
order_field=order_field,
|
||||
order=order,
|
||||
)
|
||||
return {"list": [self._param_dto(row) for row in rows], "total": total}
|
||||
|
||||
async def get(self, param_id: int) -> dict[str, Any] | None:
|
||||
row = await self.repository.get_param(param_id)
|
||||
return None if row is None else self._param_dto(row)
|
||||
|
||||
async def save(
|
||||
self,
|
||||
dto: SysParamPayload,
|
||||
user: AuthUser,
|
||||
accept_language: str | None = None,
|
||||
) -> None:
|
||||
self._validate_group(dto, update=False, accept_language=accept_language)
|
||||
self._validate_value(dto)
|
||||
assert dto.param_code is not None
|
||||
assert dto.param_value is not None
|
||||
assert dto.value_type is not None
|
||||
await self.repository.insert_param(
|
||||
param_id=snowflake.next_id(),
|
||||
param_code=dto.param_code,
|
||||
param_value=dto.param_value,
|
||||
value_type=dto.value_type,
|
||||
remark=dto.remark,
|
||||
user_id=user.id,
|
||||
now=shanghai_now_naive(),
|
||||
)
|
||||
await self._cache_set(dto.param_code, dto.param_value)
|
||||
await self.repository.session.commit()
|
||||
|
||||
async def update(
|
||||
self,
|
||||
dto: SysParamPayload,
|
||||
user: AuthUser,
|
||||
accept_language: str | None = None,
|
||||
) -> None:
|
||||
self._validate_group(dto, update=True, accept_language=accept_language)
|
||||
assert dto.id is not None
|
||||
assert dto.param_code is not None
|
||||
assert dto.param_value is not None
|
||||
assert dto.value_type is not None
|
||||
# These checks live in the Java controller and therefore run before
|
||||
# SysParamsService.update validates the declared value type.
|
||||
await self.validator.validate(dto.param_code, dto.param_value)
|
||||
if dto.param_code == "system-web.menu":
|
||||
await self._update_system_web_menu(dto.param_value, user)
|
||||
else:
|
||||
self._validate_value(dto)
|
||||
await self.repository.update_param(
|
||||
param_id=dto.id,
|
||||
param_code=dto.param_code,
|
||||
param_value=dto.param_value,
|
||||
value_type=dto.value_type,
|
||||
remark=dto.remark,
|
||||
user_id=user.id,
|
||||
now=shanghai_now_naive(),
|
||||
)
|
||||
await self._cache_set(dto.param_code, dto.param_value)
|
||||
await self.repository.session.commit()
|
||||
|
||||
async def delete(self, ids: list[str]) -> None:
|
||||
if not ids:
|
||||
raise AppError(10001, "id")
|
||||
parsed_ids = [int(value) for value in ids]
|
||||
codes = await self.repository.param_codes_for_ids(parsed_ids)
|
||||
if codes:
|
||||
await cast(Any, self.redis.hdel)("sys:params", *codes)
|
||||
await self.repository.delete_params(parsed_ids)
|
||||
await self.repository.session.commit()
|
||||
|
||||
async def get_value(self, code: str, *, from_cache: bool = True) -> str | None:
|
||||
if from_cache:
|
||||
cached = JavaRedisCodec.decode(await cast(Any, self.redis.hget)("sys:params", code))
|
||||
if cached is not None:
|
||||
return str(cached)
|
||||
value = await self.repository.get_param_value(code)
|
||||
if value is not None and from_cache:
|
||||
await self._cache_set(code, value)
|
||||
return value
|
||||
|
||||
async def config_rows(self) -> list[dict[str, Any]]:
|
||||
return await self.repository.list_config_params()
|
||||
|
||||
async def _update_system_web_menu(self, config_json: str, user: AuthUser) -> None:
|
||||
current_config = await self.repository.get_param_value("system-web.menu")
|
||||
try:
|
||||
current = json.loads(current_config) if current_config and current_config.strip() else None
|
||||
updated = json.loads(config_json) if config_json.strip() else None
|
||||
except json.JSONDecodeError as exc:
|
||||
raise AppError(ErrorCode.PARAM_JSON_INVALID) from exc
|
||||
if isinstance(current, dict) and isinstance(updated, dict):
|
||||
current_features = current.get("features")
|
||||
updated_features = updated.get("features")
|
||||
# Java only evaluates addressBook when both feature maps are present.
|
||||
if isinstance(current_features, dict) and isinstance(updated_features, dict):
|
||||
current_address = current_features.get("addressBook")
|
||||
updated_address = updated_features.get("addressBook")
|
||||
current_enabled = self._java_enabled(current_address)
|
||||
updated_enabled = self._java_enabled(updated_address)
|
||||
if current_enabled and not updated_enabled:
|
||||
await self.repository.delete_plugin_mapping_by_plugin_id("SYSTEM_PLUGIN_CALL_DEVICE")
|
||||
await self.repository.update_param_value_by_code(
|
||||
"system-web.menu", config_json, user.id, shanghai_now_naive()
|
||||
)
|
||||
await self._cache_set("system-web.menu", config_json)
|
||||
|
||||
async def _cache_set(self, code: str, value: str) -> None:
|
||||
await cast(Any, self.redis.hset)("sys:params", code, JavaRedisCodec.encode(value))
|
||||
await cast(Any, self.redis.expire)("sys:params", 24 * 60 * 60)
|
||||
|
||||
@staticmethod
|
||||
def _java_enabled(address_book: Any) -> bool:
|
||||
if not isinstance(address_book, dict):
|
||||
return False
|
||||
value = address_book.get("enabled")
|
||||
if value is None:
|
||||
return False
|
||||
if not isinstance(value, bool):
|
||||
# The Java implementation casts the JSON value to Boolean.
|
||||
raise TypeError("addressBook.enabled must be a boolean")
|
||||
return value
|
||||
|
||||
@staticmethod
|
||||
def _validate_value(dto: SysParamPayload) -> None:
|
||||
assert dto.param_value is not None
|
||||
assert dto.value_type is not None
|
||||
if not dto.param_value.strip():
|
||||
raise AppError(ErrorCode.PARAM_VALUE_NULL)
|
||||
if not dto.value_type.strip():
|
||||
raise AppError(ErrorCode.PARAM_TYPE_NULL)
|
||||
value_type = dto.value_type.lower()
|
||||
if value_type in {"string", "array"}:
|
||||
return
|
||||
if value_type == "number":
|
||||
try:
|
||||
float(dto.param_value)
|
||||
except ValueError as exc:
|
||||
raise AppError(ErrorCode.PARAM_NUMBER_INVALID) from exc
|
||||
return
|
||||
if value_type == "boolean":
|
||||
if dto.param_value.lower() not in {"true", "false"}:
|
||||
raise AppError(ErrorCode.PARAM_BOOLEAN_INVALID)
|
||||
return
|
||||
if value_type == "json":
|
||||
stripped = dto.param_value.strip()
|
||||
if not stripped.startswith("{") or not stripped.endswith("}"):
|
||||
raise AppError(ErrorCode.PARAM_JSON_INVALID)
|
||||
try:
|
||||
json.loads(dto.param_value)
|
||||
except json.JSONDecodeError as exc:
|
||||
raise AppError(ErrorCode.PARAM_JSON_INVALID) from exc
|
||||
return
|
||||
raise AppError(ErrorCode.PARAM_TYPE_INVALID)
|
||||
|
||||
@staticmethod
|
||||
def _validate_group(
|
||||
dto: SysParamPayload,
|
||||
*,
|
||||
update: bool,
|
||||
accept_language: str | None,
|
||||
) -> None:
|
||||
def fail(key: str) -> None:
|
||||
raise AppError(500, validation_message(key, accept_language))
|
||||
|
||||
if update and dto.id is None:
|
||||
fail("id.require")
|
||||
if not update and dto.id is not None:
|
||||
fail("id.null")
|
||||
if dto.param_code is None or not dto.param_code.strip():
|
||||
fail("sysparams.paramcode.require")
|
||||
if dto.param_value is None or not dto.param_value.strip():
|
||||
fail("sysparams.paramvalue.require")
|
||||
if dto.value_type is None or not dto.value_type.strip():
|
||||
fail("sysparams.valuetype.require")
|
||||
if dto.value_type not in {"string", "number", "boolean", "array", "json"}:
|
||||
fail("sysparams.valuetype.pattern")
|
||||
|
||||
@staticmethod
|
||||
def _param_dto(row: dict[str, Any]) -> dict[str, Any]:
|
||||
return {
|
||||
"id": row.get("id"),
|
||||
"paramCode": row.get("param_code"),
|
||||
"paramValue": row.get("param_value"),
|
||||
"valueType": row.get("value_type"),
|
||||
"remark": row.get("remark"),
|
||||
"createDate": row.get("create_date"),
|
||||
"updateDate": row.get("update_date"),
|
||||
}
|
||||
|
||||
|
||||
class DictService:
|
||||
def __init__(self, repository: SysRepository, *, redis: Redis | None = None):
|
||||
self.repository = repository
|
||||
self.redis = redis or get_redis()
|
||||
|
||||
async def page_types(
|
||||
self,
|
||||
*,
|
||||
dict_type: str | None,
|
||||
dict_name: str | None,
|
||||
page: int,
|
||||
limit: int,
|
||||
) -> dict[str, Any]:
|
||||
rows, total = await self.repository.page_dict_types(
|
||||
dict_type=dict_type,
|
||||
dict_name=dict_name,
|
||||
page=max(1, page),
|
||||
limit=max(0, limit),
|
||||
)
|
||||
return {"list": [self._type_vo(row, include_names=True) for row in rows], "total": total}
|
||||
|
||||
async def get_type(self, type_id: int) -> dict[str, Any]:
|
||||
row = await self.repository.get_dict_type(type_id)
|
||||
if row is None:
|
||||
raise AppError(10076)
|
||||
return self._type_vo(row, include_names=False)
|
||||
|
||||
async def save_type(self, dto: DictTypePayload, user: AuthUser) -> None:
|
||||
if await self.repository.dict_type_exists(dto.dict_type):
|
||||
raise AppError(10077)
|
||||
await self.repository.insert_dict_type(
|
||||
type_id=dto.id if dto.id is not None else snowflake.next_id(),
|
||||
dict_type=dto.dict_type,
|
||||
dict_name=dto.dict_name,
|
||||
remark=dto.remark,
|
||||
sort=dto.sort,
|
||||
user_id=user.id,
|
||||
now=shanghai_now_naive(),
|
||||
)
|
||||
await self.repository.session.commit()
|
||||
|
||||
async def update_type(self, dto: DictTypePayload, user: AuthUser) -> None:
|
||||
if await self.repository.dict_type_exists(dto.dict_type, exclude_id=dto.id):
|
||||
raise AppError(10077)
|
||||
await self.repository.update_dict_type(
|
||||
type_id=dto.id,
|
||||
dict_type=dto.dict_type,
|
||||
dict_name=dto.dict_name,
|
||||
remark=dto.remark,
|
||||
sort=dto.sort,
|
||||
user_id=user.id,
|
||||
now=shanghai_now_naive(),
|
||||
)
|
||||
await self.repository.session.commit()
|
||||
|
||||
async def delete_types(self, ids: list[int]) -> None:
|
||||
await self.repository.delete_dict_types(ids)
|
||||
await self.repository.session.commit()
|
||||
|
||||
async def page_data(
|
||||
self,
|
||||
*,
|
||||
dict_type_id: int,
|
||||
dict_label: str | None,
|
||||
dict_value: str | None,
|
||||
page: int,
|
||||
limit: int,
|
||||
) -> dict[str, Any]:
|
||||
rows, total = await self.repository.page_dict_data(
|
||||
dict_type_id=dict_type_id,
|
||||
dict_label=dict_label,
|
||||
dict_value=dict_value,
|
||||
page=max(1, page),
|
||||
limit=max(0, limit),
|
||||
)
|
||||
return {"list": [self._data_vo(row, include_names=True) for row in rows], "total": total}
|
||||
|
||||
async def get_data(self, data_id: int) -> dict[str, Any] | None:
|
||||
row = await self.repository.get_dict_data(data_id)
|
||||
return None if row is None else self._data_vo(row, include_names=False)
|
||||
|
||||
async def save_data(self, dto: DictDataPayload, user: AuthUser) -> None:
|
||||
# Java compares dict_label against dictValue here; retain that behavior for compatibility.
|
||||
if await self.repository.dict_data_label_exists(dto.dict_type_id, dto.dict_value):
|
||||
raise AppError(10128)
|
||||
await self.repository.insert_dict_data(
|
||||
data_id=dto.id if dto.id is not None else snowflake.next_id(),
|
||||
dict_type_id=dto.dict_type_id,
|
||||
dict_label=dto.dict_label,
|
||||
dict_value=dto.dict_value,
|
||||
remark=dto.remark,
|
||||
sort=dto.sort,
|
||||
user_id=user.id,
|
||||
now=shanghai_now_naive(),
|
||||
)
|
||||
await self._clear_dict_cache(dto.dict_type_id)
|
||||
await self.repository.session.commit()
|
||||
|
||||
async def update_data(self, dto: DictDataPayload, user: AuthUser) -> None:
|
||||
if await self.repository.dict_data_label_exists(dto.dict_type_id, dto.dict_value, exclude_id=dto.id):
|
||||
raise AppError(10128)
|
||||
await self.repository.update_dict_data(
|
||||
data_id=dto.id,
|
||||
dict_type_id=dto.dict_type_id,
|
||||
dict_label=dto.dict_label,
|
||||
dict_value=dto.dict_value,
|
||||
remark=dto.remark,
|
||||
sort=dto.sort,
|
||||
user_id=user.id,
|
||||
now=shanghai_now_naive(),
|
||||
)
|
||||
await self._clear_dict_cache(dto.dict_type_id)
|
||||
await self.repository.session.commit()
|
||||
|
||||
async def delete_data(self, ids: list[int]) -> None:
|
||||
if ids:
|
||||
codes = await self.repository.dict_type_codes_for_data_ids(ids)
|
||||
if codes:
|
||||
await cast(Any, self.redis.delete)(*[f"sys:dict:data:{code}" for code in codes])
|
||||
await self.repository.delete_dict_data(ids)
|
||||
await self.repository.session.commit()
|
||||
|
||||
async def items(self, dict_type: str) -> list[dict[str, Any]] | None:
|
||||
if not dict_type.strip():
|
||||
return None
|
||||
key = f"sys:dict:data:{dict_type}"
|
||||
cached = JavaRedisCodec.decode(await cast(Any, self.redis.get)(key))
|
||||
if isinstance(cached, list):
|
||||
return cast(list[dict[str, Any]], cached)
|
||||
rows = await self.repository.dict_items(dict_type)
|
||||
await cast(Any, self.redis.set)(
|
||||
key,
|
||||
JavaRedisCodec.encode(
|
||||
rows,
|
||||
item_java_type="xiaozhi.modules.sys.vo.SysDictDataItem",
|
||||
),
|
||||
ex=24 * 60 * 60,
|
||||
)
|
||||
return rows
|
||||
|
||||
async def _clear_dict_cache(self, type_id: int | None) -> None:
|
||||
dict_type = await self.repository.dict_type_code(type_id)
|
||||
if dict_type is not None:
|
||||
await cast(Any, self.redis.delete)(f"sys:dict:data:{dict_type}")
|
||||
|
||||
@staticmethod
|
||||
def _type_vo(row: dict[str, Any], *, include_names: bool) -> dict[str, Any]:
|
||||
return {
|
||||
"id": row.get("id"),
|
||||
"dictType": row.get("dict_type"),
|
||||
"dictName": row.get("dict_name"),
|
||||
"remark": row.get("remark"),
|
||||
"sort": row.get("sort"),
|
||||
"creator": row.get("creator"),
|
||||
"creatorName": row.get("creator_name") if include_names else None,
|
||||
"createDate": row.get("create_date"),
|
||||
"updater": row.get("updater"),
|
||||
"updaterName": row.get("updater_name") if include_names else None,
|
||||
"updateDate": row.get("update_date"),
|
||||
}
|
||||
|
||||
@staticmethod
|
||||
def _data_vo(row: dict[str, Any], *, include_names: bool) -> dict[str, Any]:
|
||||
return {
|
||||
"id": row.get("id"),
|
||||
"dictTypeId": row.get("dict_type_id"),
|
||||
"dictLabel": row.get("dict_label"),
|
||||
"dictValue": row.get("dict_value"),
|
||||
"remark": row.get("remark"),
|
||||
"sort": row.get("sort"),
|
||||
"creator": row.get("creator"),
|
||||
"creatorName": row.get("creator_name") if include_names else None,
|
||||
"createDate": row.get("create_date"),
|
||||
"updater": row.get("updater"),
|
||||
"updaterName": row.get("updater_name") if include_names else None,
|
||||
"updateDate": row.get("update_date"),
|
||||
}
|
||||
|
||||
|
||||
class ServerActionService:
|
||||
def __init__(self, param_service: SysParamService, *, redis: Redis | None = None):
|
||||
self.param_service = param_service
|
||||
self.redis = redis or get_redis()
|
||||
|
||||
async def server_list(self) -> list[str]:
|
||||
value = await self.param_service.get_value("server.websocket", from_cache=True)
|
||||
if value is None or not value.strip():
|
||||
return []
|
||||
values = value.split(";")
|
||||
while values and values[-1] == "":
|
||||
values.pop()
|
||||
return values
|
||||
|
||||
async def emit(self, dto: EmitServerActionRequest) -> bool:
|
||||
action = dto.action.lower() if dto.action is not None else None
|
||||
if action not in {"restart", "update_config"}:
|
||||
raise AppError(10095)
|
||||
websocket_text = await self.param_service.get_value("server.websocket", from_cache=True)
|
||||
if websocket_text is None or not websocket_text.strip():
|
||||
raise AppError(10096)
|
||||
if dto.target_ws not in websocket_text.split(";"):
|
||||
raise AppError(10097)
|
||||
payload_secret = await self.param_service.get_value("server.secret", from_cache=True)
|
||||
device_id = str(uuid.uuid4())
|
||||
client_id = str(uuid.uuid4())
|
||||
await cast(Any, self.redis.set)(
|
||||
f"tmp_register_mac:{device_id}",
|
||||
JavaRedisCodec.encode("true"),
|
||||
ex=300,
|
||||
)
|
||||
authentication_secret = await self.param_service.get_value("server.secret", from_cache=False)
|
||||
if authentication_secret is None or not authentication_secret.strip():
|
||||
raise AppError(10045)
|
||||
timestamp = int(time.time())
|
||||
content = f"{client_id}|{device_id}|{timestamp}"
|
||||
signature = hmac.new(authentication_secret.encode(), content.encode(), digestmod=hashlib.sha256).digest()
|
||||
token = base64.urlsafe_b64encode(signature).rstrip(b"=").decode() + f".{timestamp}"
|
||||
headers = {
|
||||
"device-id": device_id,
|
||||
"client-id": client_id,
|
||||
"authorization": f"Bearer {token}",
|
||||
}
|
||||
if payload_secret is None:
|
||||
raise AppError(10045)
|
||||
payload = {"type": "server", "action": action, "content": {"secret": payload_secret}}
|
||||
try:
|
||||
async with connect(dto.target_ws, additional_headers=headers, open_timeout=3) as websocket:
|
||||
await websocket.send(json.dumps(payload, ensure_ascii=False, separators=(",", ":")))
|
||||
deadline = time.monotonic() + 120
|
||||
while True:
|
||||
remaining = deadline - time.monotonic()
|
||||
if remaining <= 0:
|
||||
raise TimeoutError
|
||||
raw = await asyncio.wait_for(websocket.recv(), timeout=remaining)
|
||||
response = json.loads(raw)
|
||||
if (
|
||||
isinstance(response, dict)
|
||||
and response.get("status") == "success"
|
||||
and response.get("type") == "server"
|
||||
and isinstance(response.get("content"), dict)
|
||||
and response["content"].get("action") is not None
|
||||
):
|
||||
return True
|
||||
except Exception as exc:
|
||||
raise AppError(10045) from exc
|
||||
@@ -0,0 +1,50 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
|
||||
from sqlalchemy import text
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
|
||||
from app.core.redis import java_hget, java_hset
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
class SystemParamService:
|
||||
CACHE_KEY = "sys:params"
|
||||
|
||||
def __init__(self, session: AsyncSession):
|
||||
self.session = session
|
||||
|
||||
async def get_value(self, code: str, *, from_cache: bool = True) -> str | None:
|
||||
if from_cache:
|
||||
try:
|
||||
cached = await java_hget(self.CACHE_KEY, code)
|
||||
if cached is not None:
|
||||
return str(cached)
|
||||
except Exception:
|
||||
logger.warning("Redis parameter cache read failed for %s", code, exc_info=True)
|
||||
result = await self.session.execute(
|
||||
text("SELECT param_value FROM sys_params WHERE param_code = :code LIMIT 1"),
|
||||
{"code": code},
|
||||
)
|
||||
value = result.scalar_one_or_none()
|
||||
if value is not None and from_cache:
|
||||
try:
|
||||
await java_hset(self.CACHE_KEY, code, str(value))
|
||||
except Exception:
|
||||
logger.warning("Redis parameter cache write failed for %s", code, exc_info=True)
|
||||
return None if value is None else str(value)
|
||||
|
||||
async def set_value(self, code: str, value: str) -> int:
|
||||
result = await self.session.execute(
|
||||
text(
|
||||
"UPDATE sys_params SET param_value = :value, update_date = CURRENT_TIMESTAMP WHERE param_code = :code"
|
||||
),
|
||||
{"code": code, "value": value},
|
||||
)
|
||||
try:
|
||||
await java_hset(self.CACHE_KEY, code, value)
|
||||
except Exception:
|
||||
logger.warning("Redis parameter cache write failed for %s", code, exc_info=True)
|
||||
return int(getattr(result, "rowcount", 0) or 0)
|
||||
@@ -0,0 +1,157 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from functools import lru_cache
|
||||
from typing import Any
|
||||
|
||||
from app.core.config import get_settings
|
||||
from app.core.i18n import _load_properties, message_for, resolve_language
|
||||
from app.core.ids import snowflake
|
||||
from app.core.redis import JavaRedisCodec, get_redis
|
||||
from app.core.security import AuthUser, shanghai_now_naive
|
||||
from app.repositories.timbre import TimbreRepository
|
||||
from app.schemas.timbre import TimbreBody
|
||||
|
||||
|
||||
def _details(row: dict[str, Any]) -> dict[str, Any]:
|
||||
return {
|
||||
"id": row.get("id"),
|
||||
"languages": row.get("languages"),
|
||||
"name": row.get("name"),
|
||||
"remark": row.get("remark"),
|
||||
"referenceAudio": row.get("reference_audio"),
|
||||
"referenceText": row.get("reference_text"),
|
||||
# TimbreDetailsVO.sort is primitive long, whose Java serializer always
|
||||
# emits a string and whose null conversion default is zero.
|
||||
"sort": str(row.get("sort") if row.get("sort") is not None else 0),
|
||||
"ttsModelId": row.get("tts_model_id"),
|
||||
"ttsVoice": row.get("tts_voice"),
|
||||
"voiceDemo": row.get("voice_demo"),
|
||||
}
|
||||
|
||||
|
||||
class TimbreService:
|
||||
def __init__(self, repository: TimbreRepository):
|
||||
self.repository = repository
|
||||
|
||||
@staticmethod
|
||||
def _validate(body: TimbreBody, language: str | None) -> None:
|
||||
from app.core.errors import AppError
|
||||
|
||||
for value, message in (
|
||||
(body.languages, "timbre.languages.require"),
|
||||
(body.name, "timbre.name.require"),
|
||||
(body.tts_model_id, "timbre.ttsModelId.require"),
|
||||
(body.tts_voice, "timbre.ttsVoice.require"),
|
||||
):
|
||||
if value is None or not value.strip():
|
||||
# TimbreController invokes ValidatorUtils directly. That
|
||||
# utility wraps validation text in RenException(String), whose
|
||||
# response code is 500 rather than the global @Valid code 10034.
|
||||
raise AppError(500, _validation_message(message, language))
|
||||
if body.sort is not None and body.sort < 0:
|
||||
raise AppError(500, _validation_message("sort.number", language))
|
||||
|
||||
async def page(
|
||||
self,
|
||||
tts_model_id: str | None,
|
||||
name: str | None,
|
||||
page: str | None,
|
||||
limit: str | None,
|
||||
language: str | None,
|
||||
) -> dict[str, Any]:
|
||||
if tts_model_id is None or not tts_model_id.strip():
|
||||
from app.core.errors import AppError
|
||||
|
||||
raise AppError(500, _validation_message("timbre.ttsModelId.require", language))
|
||||
current, size = max(int(page or "1"), 1), int(limit or "10")
|
||||
rows, total = await self.repository.page(
|
||||
tts_model_id=tts_model_id, name=name, offset=(current - 1) * size, limit=size
|
||||
)
|
||||
return {"total": total, "list": [_details(row) for row in rows]}
|
||||
|
||||
async def save(self, body: TimbreBody, user: AuthUser, language: str | None) -> None:
|
||||
self._validate(body, language)
|
||||
values = self._values(body, user, str(snowflake.next_id()))
|
||||
async with self.repository.session.begin():
|
||||
await self.repository.insert(values)
|
||||
|
||||
async def update(
|
||||
self, timbre_id: str, body: TimbreBody, user: AuthUser, language: str | None
|
||||
) -> None:
|
||||
self._validate(body, language)
|
||||
values = self._values(body, user, timbre_id)
|
||||
async with self.repository.session.begin():
|
||||
await self.repository.update(values)
|
||||
await get_redis().delete(f"timbre:details:{timbre_id}")
|
||||
|
||||
async def delete(self, ids: list[str]) -> None:
|
||||
async with self.repository.session.begin():
|
||||
await self.repository.delete(ids)
|
||||
|
||||
async def voices(self, model_id: str, voice_name: str | None, user: AuthUser, language: str | None) -> Any:
|
||||
normal, clones = await self.repository.voices(model_id, voice_name, user.id)
|
||||
values = [
|
||||
{
|
||||
"id": row.get("id"),
|
||||
"name": row.get("name"),
|
||||
"voiceDemo": row.get("voice_demo"),
|
||||
"languages": row.get("languages"),
|
||||
"isClone": False,
|
||||
}
|
||||
for row in normal
|
||||
]
|
||||
prefix = message_for(10158, language)
|
||||
redis = get_redis()
|
||||
for row in clones:
|
||||
name = prefix + str(row.get("name") or "")
|
||||
voice = {
|
||||
"id": row.get("id"),
|
||||
"name": name,
|
||||
"voiceDemo": row.get("voice_demo"),
|
||||
"languages": row.get("languages"),
|
||||
"isClone": True,
|
||||
}
|
||||
await redis.set(f"timbre:name:{row['id']}", JavaRedisCodec.encode(name))
|
||||
values.insert(0, voice)
|
||||
return values or None
|
||||
|
||||
@staticmethod
|
||||
def _values(body: TimbreBody, user: AuthUser, timbre_id: str) -> dict[str, Any]:
|
||||
assert body.languages is not None
|
||||
assert body.name is not None
|
||||
assert body.tts_model_id is not None
|
||||
assert body.tts_voice is not None
|
||||
return {
|
||||
"id": timbre_id,
|
||||
"languages": body.languages,
|
||||
"name": body.name,
|
||||
"remark": body.remark,
|
||||
"reference_audio": body.reference_audio,
|
||||
"reference_text": body.reference_text,
|
||||
"sort": body.sort if body.sort is not None else 0,
|
||||
"tts_model_id": body.tts_model_id,
|
||||
"tts_voice": body.tts_voice,
|
||||
"voice_demo": body.voice_demo,
|
||||
"creator": user.id,
|
||||
"updater": user.id,
|
||||
"now": shanghai_now_naive(),
|
||||
}
|
||||
|
||||
|
||||
_VALIDATION_FILES = {
|
||||
"zh-CN": "validation_zh_CN.properties",
|
||||
"zh-TW": "validation_zh_TW.properties",
|
||||
"en-US": "validation_en_US.properties",
|
||||
"de-DE": "validation_de_DE.properties",
|
||||
"vi-VN": "validation_vi_VN.properties",
|
||||
"pt-BR": "validation_pt_BR.properties",
|
||||
}
|
||||
|
||||
|
||||
@lru_cache(maxsize=64)
|
||||
def _validation_message(key: str, accept_language: str | None) -> str:
|
||||
language = resolve_language(accept_language)
|
||||
directory = get_settings().i18n_dir
|
||||
messages = _load_properties(directory / "validation.properties")
|
||||
messages.update(_load_properties(directory / _VALIDATION_FILES[language]))
|
||||
return messages.get(key, key)
|
||||
@@ -0,0 +1,334 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
import uuid
|
||||
from collections.abc import Mapping, Sequence
|
||||
from typing import Any
|
||||
|
||||
import httpx
|
||||
from redis.asyncio import Redis
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
|
||||
from app.core.config import get_settings
|
||||
from app.core.errors import AppError
|
||||
from app.core.i18n import message_for
|
||||
from app.core.security import AuthUser, shanghai_now_naive
|
||||
from app.integrations.voice_clone import VoiceCloneIntegration, VoiceCloneProviderError
|
||||
from app.repositories.voiceclone import VoiceCloneRepository
|
||||
from app.schemas.voiceclone import VoiceResourceCreateRequest
|
||||
from app.services.device import is_blank, redis_delete, redis_get, redis_set
|
||||
|
||||
VOICE_ORDER_COLUMNS = {
|
||||
"id": "id",
|
||||
"name": "name",
|
||||
"modelId": "model_id",
|
||||
"model_id": "model_id",
|
||||
"voiceId": "voice_id",
|
||||
"voice_id": "voice_id",
|
||||
"userId": "user_id",
|
||||
"user_id": "user_id",
|
||||
"trainStatus": "train_status",
|
||||
"train_status": "train_status",
|
||||
"createDate": "create_date",
|
||||
"create_date": "create_date",
|
||||
}
|
||||
|
||||
|
||||
class VoiceCloneService:
|
||||
def __init__(
|
||||
self,
|
||||
session: AsyncSession,
|
||||
*,
|
||||
redis_client: Redis | None = None,
|
||||
http_client: httpx.AsyncClient | None = None,
|
||||
provider: VoiceCloneIntegration | None = None,
|
||||
):
|
||||
self.session = session
|
||||
self.repository = VoiceCloneRepository(session)
|
||||
self.redis = redis_client
|
||||
self.provider = provider or VoiceCloneIntegration(
|
||||
timeout_seconds=get_settings().external_request_timeout_seconds,
|
||||
client=http_client,
|
||||
)
|
||||
|
||||
async def page(self, query: Mapping[str, Any], *, user_id: int | None = None) -> dict[str, Any]:
|
||||
page = int(str(query.get("page") or "1"))
|
||||
limit = int(str(query.get("limit") or "10"))
|
||||
name_value = query.get("name")
|
||||
name = None if name_value is None else str(name_value)
|
||||
effective_user = str(user_id) if user_id is not None else self._optional_string(query.get("userId"))
|
||||
requested = query.get("orderField")
|
||||
requested_fields = [requested] if isinstance(requested, str) else list(requested or [])
|
||||
order_fields = [VOICE_ORDER_COLUMNS[field] for field in requested_fields if field in VOICE_ORDER_COLUMNS]
|
||||
if not order_fields:
|
||||
order_fields = ["create_date"]
|
||||
ascending = str(query.get("order") or "").lower() == "asc" if requested_fields else True
|
||||
rows = await self.repository.page(
|
||||
page=page,
|
||||
limit=limit,
|
||||
name=name,
|
||||
user_id=effective_user,
|
||||
order_fields=order_fields,
|
||||
ascending=ascending,
|
||||
)
|
||||
return {
|
||||
"total": await self.repository.count(name=name, user_id=effective_user),
|
||||
"list": await self._response_list(rows),
|
||||
}
|
||||
|
||||
async def get_detail(self, voice_id: str) -> dict[str, Any] | None:
|
||||
row = await self.repository.get(voice_id)
|
||||
if row is None:
|
||||
return None
|
||||
return await self._response(row, include_has_voice=False)
|
||||
|
||||
async def get_by_user(self, user_id: int) -> list[dict[str, Any]]:
|
||||
del user_id
|
||||
# VoiceCloneServiceImpl.getByUserId orders ai_voice_clone by the
|
||||
# nonexistent ``created_at`` column (the schema uses ``create_date``).
|
||||
# The Java endpoint therefore consistently exposes its generic
|
||||
# code-500 envelope before result conversion.
|
||||
raise AppError(500)
|
||||
|
||||
async def create_resources(self, request: VoiceResourceCreateRequest, *, actor: AuthUser) -> None:
|
||||
model_id = request.model_id or ""
|
||||
config = await self._model_config(model_id)
|
||||
if config is None:
|
||||
raise AppError(10152)
|
||||
provider_type = config.get("type")
|
||||
if not isinstance(provider_type, str) or not provider_type.strip():
|
||||
raise AppError(10153)
|
||||
voice_ids = request.voice_ids or []
|
||||
for voice_id in voice_ids:
|
||||
if is_blank(voice_id):
|
||||
continue
|
||||
if provider_type == "huoshan_double_stream" and "S_" not in voice_id:
|
||||
raise AppError(10160)
|
||||
if await self.repository.voice_id_count(model_id=model_id, voice_id=voice_id):
|
||||
raise AppError(10159)
|
||||
|
||||
now = shanghai_now_naive()
|
||||
prefix = now.strftime("%m%d%H%M")
|
||||
values: list[dict[str, Any]] = []
|
||||
for index, voice_id in enumerate(voice_ids, start=1):
|
||||
values.append(
|
||||
{
|
||||
"id": uuid.uuid4().hex,
|
||||
"name": f"{prefix}_{index}",
|
||||
"model_id": model_id,
|
||||
"voice_id": voice_id,
|
||||
"languages": request.languages,
|
||||
"user_id": request.user_id,
|
||||
"voice": None,
|
||||
"train_status": 0,
|
||||
"train_error": None,
|
||||
"creator": actor.id,
|
||||
"create_date": now,
|
||||
}
|
||||
)
|
||||
try:
|
||||
await self.repository.insert_many(values)
|
||||
await self.session.commit()
|
||||
except Exception:
|
||||
await self.session.rollback()
|
||||
raise
|
||||
|
||||
async def delete(self, ids: Sequence[str]) -> None:
|
||||
await self.repository.delete_many(ids)
|
||||
await self.session.commit()
|
||||
|
||||
async def check_permission(self, voice_id: str | None, user: AuthUser) -> dict[str, Any]:
|
||||
row = await self.repository.get(voice_id)
|
||||
if row is None:
|
||||
raise AppError(10144)
|
||||
if int(row.get("user_id") or -1) != user.id:
|
||||
raise AppError(10150)
|
||||
return row
|
||||
|
||||
async def upload_voice(self, voice_id: str, content: bytes) -> None:
|
||||
if await self.repository.get(voice_id) is None:
|
||||
raise AppError(10144)
|
||||
await self.repository.update_voice(voice_id, content)
|
||||
await self.session.commit()
|
||||
|
||||
async def rename(self, voice_id: str, name: str) -> None:
|
||||
if await self.repository.get(voice_id) is None:
|
||||
raise AppError(10144)
|
||||
await self.repository.update_name(voice_id, name)
|
||||
await self.session.commit()
|
||||
await redis_delete(f"timbre:name:{voice_id}", client=self.redis)
|
||||
|
||||
async def create_audio_id(self, voice_id: str) -> str:
|
||||
row = await self.repository.get(voice_id)
|
||||
if row is None or row.get("voice") is None:
|
||||
raise AppError(10182)
|
||||
value = str(uuid.uuid4())
|
||||
await redis_set(f"voiceClone:audio:id:{value}", voice_id, client=self.redis)
|
||||
return value
|
||||
|
||||
async def consume_audio(self, download_id: str) -> bytes | None:
|
||||
key = f"voiceClone:audio:id:{download_id}"
|
||||
voice_id = await redis_get(key, self.redis)
|
||||
await redis_delete(key, client=self.redis)
|
||||
if is_blank(None if voice_id is None else str(voice_id)):
|
||||
return None
|
||||
row = await self.repository.get(str(voice_id))
|
||||
data = None if row is None else row.get("voice")
|
||||
if data is None:
|
||||
return None
|
||||
result = bytes(data)
|
||||
return result or None
|
||||
|
||||
async def clone_audio(
|
||||
self,
|
||||
voice_id: str,
|
||||
*,
|
||||
accept_language: str | None,
|
||||
) -> None:
|
||||
row = await self.repository.get(voice_id)
|
||||
if row is None:
|
||||
raise AppError(10144)
|
||||
raw_voice = row.get("voice")
|
||||
if raw_voice is None or len(raw_voice) == 0:
|
||||
raise AppError(10151)
|
||||
try:
|
||||
config = await self._model_config(str(row.get("model_id") or ""))
|
||||
if config is None:
|
||||
raise AppError(10152)
|
||||
provider_type = config.get("type")
|
||||
if not isinstance(provider_type, str) or not provider_type.strip():
|
||||
raise AppError(10153)
|
||||
if provider_type != "huoshan_double_stream":
|
||||
return
|
||||
appid = config.get("appid")
|
||||
access_token = config.get("access_token")
|
||||
if (
|
||||
not isinstance(appid, str)
|
||||
or is_blank(appid)
|
||||
or not isinstance(access_token, str)
|
||||
or is_blank(access_token)
|
||||
):
|
||||
raise AppError(10155)
|
||||
speaker_id = await self.provider.train_huoshan(
|
||||
appid=appid,
|
||||
access_token=access_token,
|
||||
voice=bytes(raw_voice),
|
||||
speaker_id=str(row.get("voice_id") or ""),
|
||||
)
|
||||
await self.repository.update_training(
|
||||
voice_id,
|
||||
train_status=2,
|
||||
train_error="",
|
||||
speaker_id=speaker_id,
|
||||
)
|
||||
await self.session.commit()
|
||||
except AppError as exc:
|
||||
await self._record_training_failure(voice_id, exc.message or message_for(exc.code, accept_language))
|
||||
raise
|
||||
except VoiceCloneProviderError as exc:
|
||||
if exc.code in {500, 10156}:
|
||||
await self._record_training_failure(voice_id, exc.message)
|
||||
raise AppError(exc.code, exc.message) from exc
|
||||
translated = message_for(10154, accept_language, exc.message)
|
||||
await self._record_training_failure(voice_id, translated)
|
||||
raise AppError(10154, translated) from exc
|
||||
except Exception as exc:
|
||||
translated = message_for(10154, accept_language, str(exc))
|
||||
await self._record_training_failure(voice_id, translated)
|
||||
raise AppError(10154, translated) from exc
|
||||
|
||||
async def tts_platforms(self) -> list[dict[str, Any]]:
|
||||
return await self.repository.get_tts_platforms()
|
||||
|
||||
async def _record_training_failure(self, voice_id: str, message: str) -> None:
|
||||
await self.session.rollback()
|
||||
await self.repository.update_training(voice_id, train_status=3, train_error=message)
|
||||
await self.session.commit()
|
||||
|
||||
async def _model_config(self, model_id: str) -> dict[str, Any] | None:
|
||||
if is_blank(model_id):
|
||||
return None
|
||||
cached = await redis_get(f"model:data:{model_id}", self.redis)
|
||||
cached_mapping = self._mapping(cached)
|
||||
if cached_mapping is not None:
|
||||
config_value = cached_mapping.get("configJson", cached_mapping.get("config_json"))
|
||||
parsed = self._json_mapping(config_value)
|
||||
if parsed is not None:
|
||||
return parsed
|
||||
row = await self.repository.get_model_config(model_id)
|
||||
return None if row is None else self._json_mapping(row.get("config_json"))
|
||||
|
||||
async def _model_name(self, model_id: str | None) -> str | None:
|
||||
if is_blank(model_id):
|
||||
return None
|
||||
cache_key = f"model:name:{model_id}"
|
||||
cached = await redis_get(cache_key, self.redis)
|
||||
if isinstance(cached, str) and cached.strip():
|
||||
return cached
|
||||
value = await self.repository.get_model_name(model_id or "")
|
||||
if value is not None and value.strip():
|
||||
await redis_set(cache_key, value, client=self.redis)
|
||||
return value
|
||||
|
||||
async def _response_list(self, rows: Sequence[Mapping[str, Any]]) -> list[dict[str, Any]]:
|
||||
user_ids = [int(row["user_id"]) for row in rows if row.get("user_id") is not None]
|
||||
usernames = await self.repository.get_usernames(user_ids)
|
||||
result: list[dict[str, Any]] = []
|
||||
for row in rows:
|
||||
result.append(await self._response(row, usernames=usernames, include_has_voice=True))
|
||||
return result
|
||||
|
||||
async def _response(
|
||||
self,
|
||||
row: Mapping[str, Any],
|
||||
*,
|
||||
usernames: Mapping[int, str] | None = None,
|
||||
include_has_voice: bool,
|
||||
) -> dict[str, Any]:
|
||||
user_id = None if row.get("user_id") is None else int(row["user_id"])
|
||||
if user_id is None:
|
||||
username = None
|
||||
elif usernames is None:
|
||||
username = await self.repository.get_username(user_id)
|
||||
else:
|
||||
username = usernames.get(user_id)
|
||||
return {
|
||||
"id": row.get("id"),
|
||||
"name": row.get("name"),
|
||||
"model_id": row.get("model_id"),
|
||||
"model_name": await self._model_name(self._optional_string(row.get("model_id"))),
|
||||
"voice_id": row.get("voice_id"),
|
||||
"languages": row.get("languages"),
|
||||
"user_id": user_id,
|
||||
"user_name": username,
|
||||
"train_status": row.get("train_status"),
|
||||
"train_error": row.get("train_error"),
|
||||
"create_date": row.get("create_date"),
|
||||
"has_voice": row.get("voice") is not None if include_has_voice else None,
|
||||
}
|
||||
|
||||
@staticmethod
|
||||
def _mapping(value: Any) -> dict[str, Any] | None:
|
||||
if isinstance(value, dict):
|
||||
return {str(key): item for key, item in value.items() if key != "@class"}
|
||||
if isinstance(value, list) and len(value) == 2 and isinstance(value[1], dict):
|
||||
return {str(key): item for key, item in value[1].items()}
|
||||
return None
|
||||
|
||||
@staticmethod
|
||||
def _json_mapping(value: Any) -> dict[str, Any] | None:
|
||||
if isinstance(value, dict):
|
||||
return {str(key): item for key, item in value.items()}
|
||||
if isinstance(value, bytes):
|
||||
value = value.decode("utf-8")
|
||||
if isinstance(value, str):
|
||||
try:
|
||||
parsed = json.loads(value)
|
||||
except json.JSONDecodeError:
|
||||
return None
|
||||
return {str(key): item for key, item in parsed.items()} if isinstance(parsed, dict) else None
|
||||
return None
|
||||
|
||||
@staticmethod
|
||||
def _optional_string(value: Any) -> str | None:
|
||||
return None if value is None else str(value)
|
||||
Reference in New Issue
Block a user