init
This commit is contained in:
@@ -0,0 +1,43 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Optional
|
||||
|
||||
try:
|
||||
import redis
|
||||
except Exception:
|
||||
redis = None
|
||||
|
||||
|
||||
class CacheBase:
|
||||
"""缓存接口"""
|
||||
|
||||
def get(self, key: str) -> Optional[str]:
|
||||
raise NotImplementedError
|
||||
|
||||
def set(self, key: str, value: str, ttl: int) -> None:
|
||||
raise NotImplementedError
|
||||
|
||||
|
||||
class NoopCache(CacheBase):
|
||||
"""空实现缓存"""
|
||||
|
||||
def get(self, key: str) -> Optional[str]:
|
||||
return None
|
||||
|
||||
def set(self, key: str, value: str, ttl: int) -> None:
|
||||
return None
|
||||
|
||||
|
||||
class RedisCache(CacheBase):
|
||||
"""Redis 缓存"""
|
||||
|
||||
def __init__(self, url: str, db: int = 0):
|
||||
if redis is None:
|
||||
raise ImportError("未安装 redis 依赖")
|
||||
self._client = redis.Redis.from_url(url, db=db, decode_responses=True)
|
||||
|
||||
def get(self, key: str) -> Optional[str]:
|
||||
return self._client.get(key)
|
||||
|
||||
def set(self, key: str, value: str, ttl: int) -> None:
|
||||
self._client.set(key, value, ex=ttl)
|
||||
@@ -0,0 +1,16 @@
|
||||
from typing import Optional
|
||||
from langchain_openai import ChatOpenAI
|
||||
from config import Config
|
||||
|
||||
|
||||
def create_chat_model(model_section: Optional[str] = None) -> ChatOpenAI:
|
||||
"""创建 LLM 实例"""
|
||||
model_config = Config.get_model_config(model_section)
|
||||
return ChatOpenAI(
|
||||
model=model_config['model'],
|
||||
api_key=model_config['api_key'],
|
||||
base_url=model_config.get('base_url'),
|
||||
temperature=0.1,
|
||||
max_retries=Config.MAX_RETRIES,
|
||||
timeout=Config.TIMEOUT
|
||||
)
|
||||
@@ -0,0 +1,251 @@
|
||||
import asyncio
|
||||
import logging
|
||||
import socket
|
||||
from dataclasses import dataclass
|
||||
from typing import Any, Dict, Optional
|
||||
|
||||
try:
|
||||
import nacos
|
||||
except Exception:
|
||||
nacos = None
|
||||
|
||||
from config import Config
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
@dataclass
|
||||
class NacosConfig:
|
||||
"""Nacos 配置"""
|
||||
enabled: bool
|
||||
server_addresses: str
|
||||
namespace: str
|
||||
group_name: str
|
||||
cluster_name: str
|
||||
username: Optional[str]
|
||||
password: Optional[str]
|
||||
heartbeat_interval: int
|
||||
weight: float
|
||||
ephemeral: bool
|
||||
|
||||
|
||||
@dataclass
|
||||
class ServiceConfig:
|
||||
"""服务配置"""
|
||||
service_name: str
|
||||
host: str
|
||||
port: int
|
||||
ip: str
|
||||
metadata: Dict[str, Any]
|
||||
|
||||
|
||||
def _get_local_ip() -> str:
|
||||
"""获取本地 IP 地址"""
|
||||
try:
|
||||
s = socket.socket(socket.AF_INET, socket.SOCK_DGRAM)
|
||||
s.connect(("8.8.8.8", 80))
|
||||
ip = s.getsockname()[0]
|
||||
s.close()
|
||||
return ip
|
||||
except Exception as e:
|
||||
logger.warning(f"获取本地 IP 失败,使用 127.0.0.1: {e}")
|
||||
return "127.0.0.1"
|
||||
|
||||
|
||||
def load_nacos_config() -> NacosConfig:
|
||||
"""从 config.ini 读取 Nacos 配置"""
|
||||
section = "nacos"
|
||||
enabled = Config._config.getboolean(section, "enabled", fallback=False)
|
||||
return NacosConfig(
|
||||
enabled=enabled,
|
||||
server_addresses=Config._config.get(section, "server", fallback="localhost:8848"),
|
||||
namespace=Config._config.get(section, "namespace", fallback="public"),
|
||||
group_name=Config._config.get(section, "group_name", fallback="DEFAULT_GROUP"),
|
||||
cluster_name=Config._config.get(section, "cluster_name", fallback="DEFAULT"),
|
||||
username=Config._config.get(section, "username", fallback="") or None,
|
||||
password=Config._config.get(section, "password", fallback="") or None,
|
||||
heartbeat_interval=Config._config.getint(section, "heartbeat_interval", fallback=5),
|
||||
weight=Config._config.getfloat(section, "weight", fallback=1.0),
|
||||
ephemeral=Config._config.getboolean(section, "ephemeral", fallback=True),
|
||||
)
|
||||
|
||||
|
||||
def load_service_config() -> ServiceConfig:
|
||||
"""从 config.ini 读取服务配置"""
|
||||
section = "app"
|
||||
service_name = Config._config.get(section, "service_name", fallback="more-dots-api")
|
||||
host = Config._config.get(section, "host", fallback="0.0.0.0")
|
||||
port = Config._config.getint(section, "port", fallback=8000)
|
||||
ip = host if host != "0.0.0.0" else _get_local_ip()
|
||||
|
||||
metadata = {
|
||||
"version": Config._config.get(section, "version", fallback="1.0.0"),
|
||||
"service_type": "fastapi",
|
||||
"api_paths": "/health,/api/workflows,/api/workflows/stream,/nacos/status",
|
||||
"streaming": "false",
|
||||
"model_section": Config._config.get(section, "model_section", fallback=Config.DEFAULT_MODEL_SECTION),
|
||||
}
|
||||
|
||||
extra_meta = Config.get_section("metadata")
|
||||
if extra_meta:
|
||||
metadata.update(extra_meta)
|
||||
|
||||
return ServiceConfig(
|
||||
service_name=service_name,
|
||||
host=host,
|
||||
port=port,
|
||||
ip=ip,
|
||||
metadata=metadata,
|
||||
)
|
||||
|
||||
|
||||
class NacosManager:
|
||||
"""Nacos 服务注册与心跳管理"""
|
||||
|
||||
def __init__(self, nacos_config: NacosConfig, service_config: ServiceConfig):
|
||||
self.nacos_config = nacos_config
|
||||
self.service_config = service_config
|
||||
self.client = None
|
||||
self._heartbeat_task: Optional[asyncio.Task] = None
|
||||
self._stop_event = asyncio.Event()
|
||||
self.is_registered = False
|
||||
|
||||
def _init_client(self) -> bool:
|
||||
"""初始化 Nacos 客户端"""
|
||||
if nacos is None:
|
||||
raise ImportError("未安装 nacos-sdk-python,请先安装依赖")
|
||||
|
||||
try:
|
||||
self.client = nacos.NacosClient(
|
||||
server_addresses=self.nacos_config.server_addresses,
|
||||
namespace=self.nacos_config.namespace,
|
||||
username=self.nacos_config.username,
|
||||
password=self.nacos_config.password,
|
||||
)
|
||||
logger.info(f"Nacos 客户端初始化成功: {self.nacos_config.server_addresses}")
|
||||
return True
|
||||
except Exception as e:
|
||||
logger.error(f"Nacos 客户端初始化失败: {e}")
|
||||
return False
|
||||
|
||||
def register_service(self) -> bool:
|
||||
"""注册服务到 Nacos"""
|
||||
if not self.client and not self._init_client():
|
||||
return False
|
||||
|
||||
try:
|
||||
self.client.add_naming_instance(
|
||||
service_name=self.service_config.service_name,
|
||||
ip=self.service_config.ip,
|
||||
port=self.service_config.port,
|
||||
cluster_name=self.nacos_config.cluster_name,
|
||||
group_name=self.nacos_config.group_name,
|
||||
weight=self.nacos_config.weight,
|
||||
metadata=self.service_config.metadata,
|
||||
ephemeral=self.nacos_config.ephemeral,
|
||||
)
|
||||
self.is_registered = True
|
||||
logger.info(
|
||||
"✅ 服务注册成功: %s (%s:%s)",
|
||||
self.service_config.service_name,
|
||||
self.service_config.ip,
|
||||
self.service_config.port,
|
||||
)
|
||||
return True
|
||||
except Exception as e:
|
||||
logger.error(f"❌ 服务注册失败: {e}")
|
||||
self.is_registered = False
|
||||
return False
|
||||
|
||||
def deregister_service(self) -> bool:
|
||||
"""从 Nacos 注销服务"""
|
||||
if not self.client or not self.is_registered:
|
||||
return True
|
||||
|
||||
try:
|
||||
self.client.remove_naming_instance(
|
||||
service_name=self.service_config.service_name,
|
||||
ip=self.service_config.ip,
|
||||
port=self.service_config.port,
|
||||
cluster_name=self.nacos_config.cluster_name,
|
||||
group_name=self.nacos_config.group_name,
|
||||
)
|
||||
self.is_registered = False
|
||||
logger.info("✅ 服务注销成功: %s", self.service_config.service_name)
|
||||
return True
|
||||
except Exception as e:
|
||||
logger.error(f"❌ 服务注销失败: {e}")
|
||||
return False
|
||||
|
||||
def _send_heartbeat(self) -> None:
|
||||
"""发送心跳"""
|
||||
if not self.client or not self.is_registered:
|
||||
return
|
||||
|
||||
self.client.send_heartbeat(
|
||||
service_name=self.service_config.service_name,
|
||||
ip=self.service_config.ip,
|
||||
port=self.service_config.port,
|
||||
cluster_name=self.nacos_config.cluster_name,
|
||||
group_name=self.nacos_config.group_name,
|
||||
)
|
||||
|
||||
async def _heartbeat_loop(self) -> None:
|
||||
"""心跳循环"""
|
||||
interval = max(1, self.nacos_config.heartbeat_interval)
|
||||
self._stop_event.clear()
|
||||
|
||||
while not self._stop_event.is_set():
|
||||
try:
|
||||
self._send_heartbeat()
|
||||
logger.debug("心跳发送成功: %s", self.service_config.service_name)
|
||||
except Exception as e:
|
||||
logger.warning(f"心跳发送失败: {e}")
|
||||
# 尝试重新注册
|
||||
try:
|
||||
self.register_service()
|
||||
except Exception as re:
|
||||
logger.error(f"重新注册失败: {re}")
|
||||
|
||||
try:
|
||||
await asyncio.wait_for(self._stop_event.wait(), timeout=interval)
|
||||
except asyncio.TimeoutError:
|
||||
continue
|
||||
|
||||
async def start(self) -> None:
|
||||
"""启动注册与心跳"""
|
||||
if not self.nacos_config.enabled:
|
||||
logger.info("Nacos 未启用,跳过注册")
|
||||
return
|
||||
|
||||
if self.register_service():
|
||||
self._heartbeat_task = asyncio.create_task(self._heartbeat_loop())
|
||||
logger.info("✅ Nacos 心跳任务已启动")
|
||||
else:
|
||||
logger.warning("⚠️ Nacos 注册失败,服务继续运行")
|
||||
|
||||
async def stop(self) -> None:
|
||||
"""停止心跳并注销"""
|
||||
self._stop_event.set()
|
||||
|
||||
if self._heartbeat_task:
|
||||
self._heartbeat_task.cancel()
|
||||
try:
|
||||
await self._heartbeat_task
|
||||
except asyncio.CancelledError:
|
||||
pass
|
||||
|
||||
self.deregister_service()
|
||||
|
||||
def status(self) -> Dict[str, Any]:
|
||||
"""获取当前状态"""
|
||||
return {
|
||||
"service_name": self.service_config.service_name,
|
||||
"ip": self.service_config.ip,
|
||||
"port": self.service_config.port,
|
||||
"namespace": self.nacos_config.namespace,
|
||||
"group": self.nacos_config.group_name,
|
||||
"cluster": self.nacos_config.cluster_name,
|
||||
"registered": self.is_registered,
|
||||
"heartbeat_running": self._heartbeat_task is not None and not self._heartbeat_task.done(),
|
||||
}
|
||||
@@ -0,0 +1,31 @@
|
||||
import os
|
||||
from typing import Any, Dict, Optional
|
||||
|
||||
import yaml
|
||||
|
||||
|
||||
class PromptManager:
|
||||
"""提示词配置管理器"""
|
||||
|
||||
def __init__(self, config_path: Optional[str] = None):
|
||||
root_dir = os.path.dirname(os.path.dirname(__file__))
|
||||
self._config_path = config_path or os.path.join(root_dir, "config", "prompts.yaml")
|
||||
self._data: Dict[str, Any] = {}
|
||||
self.reload()
|
||||
|
||||
def reload(self) -> None:
|
||||
"""重新加载提示词配置"""
|
||||
with open(self._config_path, "r", encoding="utf-8") as f:
|
||||
self._data = yaml.safe_load(f) or {}
|
||||
|
||||
def get(self, group: str, name: str, default: str = "") -> str:
|
||||
"""获取指定提示词"""
|
||||
return str(self._data.get(group, {}).get(name, default))
|
||||
|
||||
def list_groups(self) -> list[str]:
|
||||
"""列出所有分组"""
|
||||
return list(self._data.keys())
|
||||
|
||||
def list_prompts(self, group: str) -> list[str]:
|
||||
"""列出分组内提示词"""
|
||||
return list(self._data.get(group, {}).keys())
|
||||
@@ -0,0 +1,59 @@
|
||||
import json
|
||||
from typing import Any, Dict, Optional
|
||||
|
||||
import httpx
|
||||
|
||||
from config import Config
|
||||
|
||||
|
||||
class RagflowClient:
|
||||
"""RAGFlow 客户端(仅检索)"""
|
||||
|
||||
def __init__(self):
|
||||
cfg = Config.get_section("ragflow")
|
||||
self._base_url = cfg.get("url", "")
|
||||
self._api_key = cfg.get("api_key", "")
|
||||
self._retrieval_path = cfg.get("retrieval", "/api/v1/retrieval")
|
||||
self._dataset_ids = cfg.get("dataset_ids", "")
|
||||
|
||||
def _build_url(self) -> str:
|
||||
return self._base_url.rstrip("/") + "/" + self._retrieval_path.lstrip("/")
|
||||
|
||||
def retrieve(self, query: str, top_k: int = 3) -> Dict[str, Any]:
|
||||
"""检索匹配文档"""
|
||||
if not self._base_url or not self._retrieval_path:
|
||||
raise RuntimeError("未配置 ragflow.url 或 ragflow.retrieval")
|
||||
url = self._build_url()
|
||||
headers = {"Authorization": f"Bearer {self._api_key}"} if self._api_key else {}
|
||||
payload = {
|
||||
"dataset_ids": self._dataset_ids,
|
||||
"query": query,
|
||||
"top_k": top_k,
|
||||
}
|
||||
|
||||
with httpx.Client(timeout=30) as client:
|
||||
response = client.post(url, json=payload, headers=headers)
|
||||
response.raise_for_status()
|
||||
return response.json()
|
||||
|
||||
|
||||
def extract_table_name(record: Dict[str, Any]) -> Optional[str]:
|
||||
"""从检索结果中提取表名"""
|
||||
if not record:
|
||||
return None
|
||||
|
||||
metadata = record.get("metadata") or {}
|
||||
for key in ("table", "table_name"):
|
||||
if key in metadata:
|
||||
return metadata.get(key)
|
||||
|
||||
for key in ("table", "table_name"):
|
||||
if key in record:
|
||||
return record.get(key)
|
||||
|
||||
content = record.get("content") or record.get("text") or ""
|
||||
for line in str(content).splitlines():
|
||||
if line.lower().startswith("table:"):
|
||||
return line.split(":", 1)[1].strip()
|
||||
|
||||
return None
|
||||
@@ -0,0 +1,64 @@
|
||||
import hashlib
|
||||
import json
|
||||
from typing import Any, Dict, Optional
|
||||
|
||||
from config import Config
|
||||
from services.cache import NoopCache, RedisCache
|
||||
from services.ragflow_client import RagflowClient, extract_table_name
|
||||
|
||||
|
||||
class TemplateMatcher:
|
||||
"""模板匹配器:RAGFlow + Redis 缓存"""
|
||||
|
||||
def __init__(self):
|
||||
self._ragflow = RagflowClient()
|
||||
self._cache = self._init_cache()
|
||||
cfg = Config.get_section("ragflow")
|
||||
self._cache_ttl = int(cfg.get("cache_ttl", 600))
|
||||
|
||||
def _init_cache(self):
|
||||
cfg = Config.get_section("redis")
|
||||
enabled = str(cfg.get("enabled", "false")).lower() in ("1", "true", "yes")
|
||||
if not enabled:
|
||||
return NoopCache()
|
||||
|
||||
url = cfg.get("url")
|
||||
db = int(cfg.get("db", 0))
|
||||
if not url:
|
||||
return NoopCache()
|
||||
|
||||
try:
|
||||
return RedisCache(url=url, db=db)
|
||||
except Exception:
|
||||
return NoopCache()
|
||||
|
||||
@staticmethod
|
||||
def _cache_key(text: str) -> str:
|
||||
return "ragflow:table:" + hashlib.sha256(text.encode("utf-8")).hexdigest()
|
||||
|
||||
def match(self, normalized_text: str) -> Dict[str, Any]:
|
||||
"""返回匹配的表名与原始响应"""
|
||||
key = self._cache_key(normalized_text)
|
||||
cached = self._cache.get(key)
|
||||
if cached:
|
||||
return json.loads(cached)
|
||||
try:
|
||||
response = self._ragflow.retrieve(normalized_text, top_k=3)
|
||||
except Exception as e:
|
||||
result = {"table_name": None, "raw": {"error": str(e)}}
|
||||
self._cache.set(key, json.dumps(result, ensure_ascii=False), self._cache_ttl)
|
||||
return result
|
||||
|
||||
candidates = []
|
||||
data = response.get("data") if isinstance(response, dict) else None
|
||||
if isinstance(data, list):
|
||||
for item in data:
|
||||
table_name = extract_table_name(item)
|
||||
if table_name:
|
||||
candidates.append(table_name)
|
||||
|
||||
matched = candidates[0] if candidates else None
|
||||
result = {"table_name": matched, "raw": response}
|
||||
|
||||
self._cache.set(key, json.dumps(result, ensure_ascii=False), self._cache_ttl)
|
||||
return result
|
||||
@@ -0,0 +1,41 @@
|
||||
import json
|
||||
from typing import Any, Dict, Optional
|
||||
|
||||
from langchain_core.tools import BaseTool
|
||||
|
||||
from tools.calculator import CalculatorTool
|
||||
from tools.web_search import WebSearchTool
|
||||
from tools.rest_api_tool import RestApiTool
|
||||
from tools.sr_api_tool import SrApiQueryTool
|
||||
|
||||
|
||||
class ToolRouter:
|
||||
"""工具路由器:统一调用入口"""
|
||||
|
||||
def __init__(self, tools: Optional[list[BaseTool]] = None):
|
||||
if tools is None:
|
||||
tools = [CalculatorTool(), WebSearchTool(), RestApiTool(), SrApiQueryTool()]
|
||||
self._tools: Dict[str, BaseTool] = {tool.name: tool for tool in tools}
|
||||
|
||||
def list_tools(self) -> list[str]:
|
||||
"""列出可用工具名称"""
|
||||
return list(self._tools.keys())
|
||||
|
||||
def call(self, tool_name: str, payload: Any) -> Dict[str, Any]:
|
||||
"""调用工具并返回标准化结果"""
|
||||
tool = self._tools.get(tool_name)
|
||||
if not tool:
|
||||
return {"ok": False, "data": None, "error": f"工具不存在: {tool_name}"}
|
||||
|
||||
try:
|
||||
if isinstance(payload, (dict, list)):
|
||||
input_value = json.dumps(payload, ensure_ascii=False)
|
||||
elif payload is None:
|
||||
input_value = ""
|
||||
else:
|
||||
input_value = str(payload)
|
||||
|
||||
result = tool.run(input_value)
|
||||
return {"ok": True, "data": result, "error": None}
|
||||
except Exception as e:
|
||||
return {"ok": False, "data": None, "error": str(e)}
|
||||
Reference in New Issue
Block a user