chore: snapshot in-progress Python-to-Rust service migration
Preserves the device's working-tree progress in Gitea as the canonical source. This is an intermediate state of an unfinished refactor, kept so the work is not lost and the merge into KVM can build on a known point. - Python services removed (superseded by Rust): privacy_gateway, rkllm_server, mem_bridge. - Rust services added/updated: privacy-gateway-rs, info-privacy-rs, kvm-agent-rs (new api-server and runner crates), embed-db-rs, npu_daemon (new rga / rga-sys / mpp-jpeg crates), rkllm-server. - Debian packaging restructured: per-service control / postinst / postrm / conffiles and systemd units. Compiled binaries (debian/*/usr/sbin), NPU model blobs and test-image fixtures are excluded from version control - see .gitignore. Co-Authored-By: Claude Opus 4.7 (1M context) <noreply@anthropic.com>
This commit is contained in:
@@ -35,3 +35,9 @@ testdata/results/
|
||||
# IDE / tooling
|
||||
.playwright-mcp/
|
||||
firebase-debug.log
|
||||
|
||||
# Build artifacts — compiled binaries and model blobs (reproducible from source)
|
||||
debian/*/usr/sbin/
|
||||
debian/kvm-npu/usr/share/
|
||||
testdata/
|
||||
tools/workflow-dashboard/testdata/
|
||||
|
||||
@@ -13,13 +13,12 @@ KVM-Privacy 是一个基于 KVM-over-IP 的两层隐私保护系统,运行在
|
||||
| 服务 | 端口 | DEB 包 | 说明 |
|
||||
|------|------|--------|------|
|
||||
| KVM Server (Go) | 8080 | kvm-server | KVM 控制 + React WebUI + WebRTC |
|
||||
| info-privacy-rs (Rust) | 8001 | kvm-privacy | RKNN PII 检测/脱敏 |
|
||||
| info-privacy-rs (Rust) | 8001 | kvm-privacy | RKNN PII 检测/脱敏 (NPU Bridge only, no Python workers) |
|
||||
| NPU Daemon (Rust) | 8004 | kvm-npu | 集中 RKNN 推理 (OCR/Face) |
|
||||
| mem-bridge memory | 8001 | kvm-bridge | 会话存储 + FAISS 向量搜索 |
|
||||
| mem-bridge router | 8002 | kvm-bridge | AI 路由服务 |
|
||||
| Privacy Gateway (Python) | 8888 | kvm-mitm | mitmproxy 网络拦截 |
|
||||
| KVM Agent (Python) | 8890 | kvm-agent | AI Agent daemon |
|
||||
| RKLLM Server (Python) | 8891 | kvm-rkllm | 本地 LLM (Qwen3-0.6B NPU) |
|
||||
| embed-db (Rust) | 8003+8002 | kvm-bridge | USearch 向量搜索 + SQLite 会话存储 + AI 路由 |
|
||||
| Privacy Gateway (Rust) | 8888+8889 | kvm-mitm | hudsucker 透明代理 + axum API |
|
||||
| KVM Agent (Python) | 8890 | kvm-agent | AI Agent daemon (pending Rust migration) |
|
||||
| RKLLM Server (Rust) | 8891 | kvm-rkllm | 本地 LLM (Qwen3-0.6B NPU, libloading FFI) |
|
||||
|
||||
### 架构
|
||||
|
||||
@@ -39,25 +38,27 @@ KVM-Privacy 是一个基于 KVM-over-IP 的两层隐私保护系统,运行在
|
||||
|
||||
```
|
||||
services/
|
||||
kvm_agent/ # Python - KVM AI Agent
|
||||
privacy_gateway/ # Python - mitmproxy 隐私网关
|
||||
rkllm_server/ # Python - RKLLM 本地 LLM 服务
|
||||
kvm_agent/ # Python - KVM AI Agent (pending Rust migration P3/P4)
|
||||
privacy-gateway-rs/ # Rust - hudsucker 隐私网关 (replaces Python mitmproxy)
|
||||
rkllm-server/ # Rust - RKLLM NPU LLM 服务 (replaces Python FastAPI)
|
||||
embed-db-rs/ # Rust - 向量搜索 + 会话存储 (replaces Python mem-bridge)
|
||||
npu_daemon/ # Rust - NPU 推理守护进程
|
||||
tools/
|
||||
workflow-dashboard/ # 测试工具(非生产),独立 pyproject.toml
|
||||
smoke-test.sh # 服务健康检查脚本 (replaces workflow-dashboard)
|
||||
KVM/ # Git submodule - Go KVM 服务端
|
||||
deps/ # Git submodules - 依赖项目
|
||||
deploy/systemd/ # systemd 服务文件
|
||||
debian/ # DEB 包定义 (kvm-mitm, kvm-agent, kvm-bridge, kvm-npu, kvm-privacy, kvm-rkllm, kvm-meta)
|
||||
scripts/build-debs.sh # 全栈 DEB 构建脚本
|
||||
scripts/build-debs.sh # 全栈 DEB 构建脚本 (Rust + Python agent)
|
||||
```
|
||||
|
||||
## 编码规范
|
||||
- Python:asyncio + httpx,类型注解,dataclass 优先
|
||||
- 异步优先:所有 I/O 操作使用 async/await
|
||||
- Rust:tokio + axum + reqwest,serde 序列化,thiserror 错误处理
|
||||
- Python(仅 kvm_agent,待迁移):asyncio + httpx,类型注解,dataclass 优先
|
||||
- 异步优先:所有 I/O 操作使用 async/await (Rust: tokio, Python: asyncio)
|
||||
- 错误处理:finally 块中释放资源(HID 控制、隐私模式)
|
||||
- 测试:pytest + pytest-asyncio,mock 外部依赖
|
||||
- 配置:YAML 文件 + 环境变量,dataclass 承载
|
||||
- 测试:Rust #[tokio::test],Python pytest + pytest-asyncio,mock 外部依赖
|
||||
- 配置:YAML 文件 + 环境变量
|
||||
|
||||
## 测试方法论
|
||||
|
||||
@@ -146,7 +147,7 @@ cd deps/KVM/go && /usr/local/go/bin/go build ./...
|
||||
# kvm-server deb 打包(需要 Go 1.24 在 PATH 前面)
|
||||
cd deps/KVM && PATH="/usr/local/go/bin:$PATH" bash scripts/build-deb.sh
|
||||
|
||||
# 构建全部 DEB 包(Python + Rust + 闭源库)
|
||||
# 构建全部 DEB 包(Rust + Python agent + 闭源库)
|
||||
bash scripts/build-debs.sh
|
||||
# 跳过 Rust 编译(使用已有二进制)
|
||||
bash scripts/build-debs.sh --skip-rust
|
||||
@@ -154,23 +155,43 @@ bash scripts/build-debs.sh --skip-rust
|
||||
# 前端编译
|
||||
cd deps/KVM/web && npm run build
|
||||
|
||||
# Dashboard (测试工具,非生产)
|
||||
cd tools/workflow-dashboard && python3 -m workflow_dashboard
|
||||
# Dashboard (Docker)
|
||||
cd tools/workflow-dashboard && docker compose up -d
|
||||
# 服务健康检查
|
||||
bash tools/smoke-test.sh
|
||||
|
||||
# 查看服务状态
|
||||
systemctl status kvm-agent mem-bridge-memory mem-bridge-router info-privacy privacy-gateway npu-daemon rkllm-server
|
||||
systemctl status kvm-agent embed-db info-privacy kvm-mitm npu-daemon rkllm-server
|
||||
|
||||
# 设备连接
|
||||
ssh pi@192.168.123.181
|
||||
```
|
||||
|
||||
## 部署规则(强制)
|
||||
|
||||
### 禁止直接拷贝二进制文件部署
|
||||
|
||||
- **绝对禁止** `cp`/`scp` 二进制文件到 `/usr/bin/`、`/usr/sbin/` 等系统路径
|
||||
- **所有服务部署必须通过 DEB 包**:`scripts/build-debs.sh` → `dpkg -i`
|
||||
- 原因:直接拷贝绕过 systemd 服务管理、conffile 保护、依赖检查、卸载清理
|
||||
- 唯一例外:开发时 `cargo run` / `python -m` 本地调试(不部署到系统路径)
|
||||
|
||||
### 标准部署流程
|
||||
|
||||
```bash
|
||||
# 1. 编译 + 打包(全栈)
|
||||
bash scripts/build-debs.sh
|
||||
|
||||
# 2. 安装指定包
|
||||
sudo dpkg -i debian/dist/kvm-privacy_*.deb
|
||||
|
||||
# 3. 验证服务状态
|
||||
systemctl is-active info-privacy && echo "PASS" || echo "FAIL"
|
||||
```
|
||||
|
||||
## Deb 包开发完整工作流
|
||||
|
||||
### 修改 → 构建 → 验证
|
||||
|
||||
每次修改 deb 包相关文件(debian/control, postinst, prerm, systemd service, C 源码)后,按以下流程验证:
|
||||
每次修改 deb 包相关文件(debian/control, postinst, prerm, systemd service, Rust/Go 源码)后,按以下流程验证:
|
||||
|
||||
```bash
|
||||
# 1. 构建
|
||||
@@ -241,11 +262,13 @@ sudo dpkg -i /tmp/kvm-server_*.deb
|
||||
| 过时 timer 噪音 | kvm-ocr-snapshot.timer 每秒触发 oneshot,native 模式不需要 | 从包中移除 |
|
||||
|
||||
## 目标设备
|
||||
- 硬件:NanoPC-T6 (RK3588, 8核 ARM64, 6 TOPS NPU)
|
||||
- 硬件:NanoPC-T6 (RK3588, 8核 ARM64, 8GB RAM, 6 TOPS NPU)
|
||||
- 开发与运行同机:代码编辑、编译、服务运行均在本机完成,资源共享
|
||||
- 系统:Debian/Ubuntu ARM64
|
||||
- Python:3.12.3
|
||||
- Go:1.22+
|
||||
- NPU:rknn-toolkit-lite2 2.3.2
|
||||
- 资源约束:8GB RAM 需在多服务间分配,CPU 大核(A76)给 Go+HID,小核(A55)给 PII+mem-bridge
|
||||
|
||||
## 已知限制
|
||||
|
||||
@@ -254,7 +277,7 @@ sudo dpkg -i /tmp/kvm-server_*.deb
|
||||
| 安全模型 | ✅ 良好 | 两层防护(视频遮蔽+网络拦截),白名单模式,Unicode NFKC 归一化 |
|
||||
| 鼠标优先 | ✅ 完成 | mouse_ops + LLM 提示词 + agent 翻译层 |
|
||||
| 多 UI 状态 | ✅ 完成 | screen_state.py 检测 BIOS/锁屏/睡眠/桌面 |
|
||||
| 测试覆盖 | ✅ 良好 | 单元测试 331 个;integration 需实体设备(--ignore 跳过) |
|
||||
| 测试覆盖 | ✅ 良好 | 单元测试 462 个;integration 需实体设备(--ignore 跳过) |
|
||||
| 隐私遮蔽 | ✅ 完成 | C 视频管道实时 NV12 黑色填充 PII 区域 (WF6 Phase 6A-2) |
|
||||
| test_integration 一致性 | ⚠️ 部分 | 清理代码已迁移到 mouse_ops,测试目标仍用原始组合键 |
|
||||
| Privacy/Audit 重组 | ✅ 完成 | PrivacyPage 5 Tab,AuditPage 2 Tab,新增 5 条 API 路由 |
|
||||
|
||||
Vendored
+1
-3
@@ -7,10 +7,7 @@ case "$1" in
|
||||
cat > /etc/kvm-agent/secrets.env <<'EOF'
|
||||
KVM_JWT_TOKEN=
|
||||
LLM_API_KEY=
|
||||
KVM_AGENT_DB_HOST=localhost
|
||||
KVM_AGENT_DB_USER=kvm_agent
|
||||
KVM_AGENT_DB_PASS=changeme
|
||||
KVM_AGENT_DB_NAME=kvm
|
||||
EOF
|
||||
chmod 600 /etc/kvm-agent/secrets.env
|
||||
chown root:root /etc/kvm-agent/secrets.env
|
||||
@@ -19,6 +16,7 @@ EOF
|
||||
|
||||
systemctl daemon-reload
|
||||
systemctl enable kvm-agent.service 2>/dev/null || true
|
||||
systemctl start kvm-agent.service 2>/dev/null || true
|
||||
;;
|
||||
esac
|
||||
exit 0
|
||||
|
||||
Vendored
+2
-5
@@ -2,12 +2,9 @@
|
||||
set -e
|
||||
case "$1" in
|
||||
purge)
|
||||
# Clean runtime-generated __pycache__ files
|
||||
rm -rf /usr/lib/kvm-agent/kvm_agent/__pycache__
|
||||
rmdir --ignore-fail-on-non-empty /usr/lib/kvm-agent/kvm_agent 2>/dev/null || true
|
||||
rmdir --ignore-fail-on-non-empty /usr/lib/kvm-agent 2>/dev/null || true
|
||||
# Clean config directory
|
||||
rm -rf /etc/kvm-agent
|
||||
rm -rf /run/kvm-agent
|
||||
rm -rf /var/lib/kvm-agent
|
||||
;;
|
||||
esac
|
||||
exit 0
|
||||
|
||||
+1
-1
@@ -1,6 +1,6 @@
|
||||
[Unit]
|
||||
Description=KVM AI Agent v2
|
||||
After=network.target mariadb.service
|
||||
After=network.target embed-db.service
|
||||
Wants=embed-db.service
|
||||
StartLimitInterval=300
|
||||
StartLimitBurst=5
|
||||
|
||||
Vendored
-3
@@ -1,3 +0,0 @@
|
||||
#!/bin/bash
|
||||
export PYTHONPATH=/usr/lib/kvm-agent
|
||||
exec python3 -m kvm_agent "$@"
|
||||
+1
@@ -0,0 +1 @@
|
||||
/usr/sbin/kvm-agent
|
||||
Vendored
+1
@@ -0,0 +1 @@
|
||||
/etc/kvm-bridge/bridge.env
|
||||
Vendored
+6
-7
@@ -1,9 +1,8 @@
|
||||
Package: kvm-bridge
|
||||
Version: 1.0.0-1
|
||||
Architecture: all
|
||||
Version: 2.0.0-1
|
||||
Architecture: arm64
|
||||
Maintainer: KVM-Privacy <noreply@kvm-privacy.local>
|
||||
Depends: python3 (>= 3.9), python3-venv
|
||||
Description: KVM mem-bridge session memory service
|
||||
FAISS vector search + SQLite session storage for AI agent memory.
|
||||
Provides memory and router sub-services.
|
||||
Installed in isolated Python venv at /opt/kvm-bridge.
|
||||
Depends: libc6
|
||||
Description: KVM embed-db session memory service (Rust)
|
||||
USearch vector search + SQLite session storage for AI agent memory.
|
||||
Single binary serving memory API (8003) and router API (8002).
|
||||
|
||||
Vendored
+13
-14
@@ -1,23 +1,22 @@
|
||||
#!/bin/bash
|
||||
set -e
|
||||
|
||||
VENV_DIR="/opt/kvm-bridge"
|
||||
LIB_DIR="/usr/lib/kvm-bridge"
|
||||
REQ_FILE="$LIB_DIR/requirements.txt"
|
||||
|
||||
if [ "$1" = "configure" ]; then
|
||||
# Create venv if not exists
|
||||
if [ ! -d "$VENV_DIR" ]; then
|
||||
python3 -m venv "$VENV_DIR"
|
||||
# Create config directory and default env file
|
||||
mkdir -p /etc/kvm-bridge
|
||||
if [ ! -f /etc/kvm-bridge/bridge.env ]; then
|
||||
cat > /etc/kvm-bridge/bridge.env <<'EOF'
|
||||
BRIDGE_DATA_DIR=/var/lib/kvm-bridge
|
||||
EMBED_MODEL_PATH=/usr/share/kvm-bridge/models/embedder.onnx
|
||||
EOF
|
||||
chmod 644 /etc/kvm-bridge/bridge.env
|
||||
fi
|
||||
|
||||
# Install dependencies
|
||||
if [ -f "$REQ_FILE" ]; then
|
||||
"$VENV_DIR/bin/pip" install --quiet -r "$REQ_FILE" || true
|
||||
fi
|
||||
# Create runtime data directory
|
||||
mkdir -p /var/lib/kvm-bridge
|
||||
chown pi:pi /var/lib/kvm-bridge
|
||||
|
||||
systemctl daemon-reload
|
||||
systemctl enable mem-bridge-memory.service mem-bridge-router.service || true
|
||||
# Do not auto-start — bridge is disabled by default in kvm.toml
|
||||
echo "kvm-bridge installed. Enable in /etc/kvm/kvm.toml: [services.bridge] enabled = true"
|
||||
systemctl enable embed-db.service || true
|
||||
systemctl start embed-db.service 2>/dev/null || true
|
||||
fi
|
||||
|
||||
+9
@@ -0,0 +1,9 @@
|
||||
#!/bin/bash
|
||||
set -e
|
||||
case "$1" in
|
||||
purge)
|
||||
rm -rf /etc/kvm-bridge
|
||||
rm -rf /var/lib/kvm-bridge
|
||||
;;
|
||||
esac
|
||||
exit 0
|
||||
Vendored
+7
-9
@@ -1,11 +1,9 @@
|
||||
#!/bin/bash
|
||||
set -e
|
||||
|
||||
if [ "$1" = "remove" ] || [ "$1" = "purge" ]; then
|
||||
systemctl stop mem-bridge-memory.service mem-bridge-router.service || true
|
||||
systemctl disable mem-bridge-memory.service mem-bridge-router.service || true
|
||||
fi
|
||||
|
||||
if [ "$1" = "purge" ]; then
|
||||
rm -rf /opt/kvm-bridge
|
||||
fi
|
||||
case "$1" in
|
||||
remove|purge)
|
||||
systemctl stop embed-db.service 2>/dev/null || true
|
||||
systemctl disable embed-db.service 2>/dev/null || true
|
||||
;;
|
||||
esac
|
||||
exit 0
|
||||
|
||||
+2
@@ -0,0 +1,2 @@
|
||||
BRIDGE_DATA_DIR=/var/lib/kvm-bridge
|
||||
EMBED_MODEL_PATH=/usr/share/kvm-bridge/models/embedder.onnx
|
||||
@@ -0,0 +1,28 @@
|
||||
[Unit]
|
||||
Description=Embed-DB Memory + Router Service (Rust, ports 8003+8002)
|
||||
After=network.target
|
||||
Wants=network.target
|
||||
StartLimitInterval=300
|
||||
StartLimitBurst=5
|
||||
|
||||
[Service]
|
||||
Type=simple
|
||||
User=pi
|
||||
Group=pi
|
||||
EnvironmentFile=-/etc/kvm-bridge/bridge.env
|
||||
Environment=MEMORY_PORT=8003
|
||||
Environment=ROUTER_PORT=8002
|
||||
ExecStart=/usr/sbin/embed-db
|
||||
Restart=on-failure
|
||||
RestartSec=10
|
||||
MemoryMax=1G
|
||||
CPUAffinity=4 5 6 7
|
||||
LimitNOFILE=4096
|
||||
StateDirectory=kvm-bridge
|
||||
RuntimeDirectory=kvm-bridge
|
||||
RuntimeDirectoryMode=0750
|
||||
StandardOutput=journal
|
||||
StandardError=journal
|
||||
|
||||
[Install]
|
||||
WantedBy=multi-user.target
|
||||
@@ -1,26 +0,0 @@
|
||||
[Unit]
|
||||
Description=Mem-Bridge Memory Service
|
||||
After=network.target
|
||||
Wants=network.target
|
||||
StartLimitInterval=300
|
||||
StartLimitBurst=5
|
||||
|
||||
[Service]
|
||||
Type=simple
|
||||
User=pi
|
||||
Group=pi
|
||||
WorkingDirectory=/home/pi/Desktop/embed-db
|
||||
Environment=PYTHONPATH=/home/pi/Desktop/embed-db/src
|
||||
EnvironmentFile=-/home/pi/Desktop/embed-db/.env
|
||||
ExecStart=/home/pi/Desktop/embed-db/venv/bin/python server.py --service memory
|
||||
Restart=on-failure
|
||||
RestartSec=10
|
||||
MemoryMax=512M
|
||||
# RK3588: pin to A55 small cores 2-3 (Python FAISS + embedding)
|
||||
CPUAffinity=2 3
|
||||
LimitNOFILE=4096
|
||||
StandardOutput=journal
|
||||
StandardError=journal
|
||||
|
||||
[Install]
|
||||
WantedBy=multi-user.target
|
||||
@@ -1,27 +0,0 @@
|
||||
[Unit]
|
||||
Description=Mem-Bridge Router Service
|
||||
After=network.target mem-bridge-memory.service
|
||||
Wants=network.target
|
||||
Requires=mem-bridge-memory.service
|
||||
StartLimitInterval=300
|
||||
StartLimitBurst=5
|
||||
|
||||
[Service]
|
||||
Type=simple
|
||||
User=pi
|
||||
Group=pi
|
||||
WorkingDirectory=/home/pi/Desktop/embed-db
|
||||
Environment=PYTHONPATH=/home/pi/Desktop/embed-db/src
|
||||
EnvironmentFile=-/home/pi/Desktop/embed-db/.env
|
||||
ExecStart=/home/pi/Desktop/embed-db/venv/bin/python server.py --service router
|
||||
Restart=on-failure
|
||||
RestartSec=10
|
||||
MemoryMax=2G
|
||||
# RK3588: pin to A55 small cores 2-3 (AI routing + LLM orchestration)
|
||||
CPUAffinity=2 3
|
||||
LimitNOFILE=4096
|
||||
StandardOutput=journal
|
||||
StandardError=journal
|
||||
|
||||
[Install]
|
||||
WantedBy=multi-user.target
|
||||
-3
@@ -1,3 +0,0 @@
|
||||
#!/bin/bash
|
||||
export PYTHONPATH=/usr/lib/kvm-bridge
|
||||
exec /opt/kvm-bridge/bin/python /usr/lib/kvm-bridge/server.py --service memory "$@"
|
||||
-3
@@ -1,3 +0,0 @@
|
||||
#!/bin/bash
|
||||
export PYTHONPATH=/usr/lib/kvm-bridge
|
||||
exec /opt/kvm-bridge/bin/python /usr/lib/kvm-bridge/server.py --service router "$@"
|
||||
@@ -1 +0,0 @@
|
||||
"""mem-bridge: 记忆注入 + 多模型路由中间层。"""
|
||||
@@ -1,95 +0,0 @@
|
||||
"""后端适配器:统一 chat() 接口,支持 Anthropic / OpenAI / 兼容格式。"""
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
from typing import AsyncIterator
|
||||
|
||||
from mem_bridge.config import BackendConfig
|
||||
from mem_bridge.models import ChatMessage
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
class BaseBackend:
|
||||
"""所有后端的基类。"""
|
||||
async def chat(
|
||||
self,
|
||||
messages: list[ChatMessage],
|
||||
stream: bool = False,
|
||||
temperature: float | None = None,
|
||||
max_tokens: int | None = None,
|
||||
) -> str | AsyncIterator[str]:
|
||||
raise NotImplementedError
|
||||
|
||||
|
||||
class AnthropicBackend(BaseBackend):
|
||||
def __init__(self, cfg: BackendConfig) -> None:
|
||||
import anthropic
|
||||
self._client = anthropic.AsyncAnthropic(api_key=cfg.api_key)
|
||||
self._model = cfg.model
|
||||
|
||||
async def chat(self, messages, stream=False, temperature=None, max_tokens=None):
|
||||
system_parts = [m.content for m in messages if m.role == "system"]
|
||||
user_messages = [
|
||||
{"role": m.role, "content": m.content}
|
||||
for m in messages if m.role != "system"
|
||||
]
|
||||
kwargs: dict = {"model": self._model, "messages": user_messages, "max_tokens": max_tokens or 2048}
|
||||
if system_parts:
|
||||
kwargs["system"] = "\n".join(system_parts)
|
||||
if temperature is not None:
|
||||
kwargs["temperature"] = temperature
|
||||
|
||||
if stream:
|
||||
async def _stream_gen():
|
||||
async with self._client.messages.stream(**kwargs) as s:
|
||||
async for text in s.text_stream:
|
||||
yield text
|
||||
return _stream_gen()
|
||||
|
||||
resp = await self._client.messages.create(**kwargs)
|
||||
return resp.content[0].text
|
||||
|
||||
|
||||
class OpenAIBackend(BaseBackend):
|
||||
def __init__(self, cfg: BackendConfig) -> None:
|
||||
from openai import AsyncOpenAI
|
||||
kwargs: dict = {"api_key": cfg.api_key}
|
||||
if cfg.base_url:
|
||||
kwargs["base_url"] = cfg.base_url
|
||||
self._client = AsyncOpenAI(**kwargs)
|
||||
self._model = cfg.model
|
||||
|
||||
async def chat(self, messages, stream=False, temperature=None, max_tokens=None):
|
||||
payload = [{"role": m.role, "content": m.content} for m in messages]
|
||||
kwargs: dict = {"model": self._model, "messages": payload, "stream": stream}
|
||||
if temperature is not None:
|
||||
kwargs["temperature"] = temperature
|
||||
if max_tokens is not None:
|
||||
kwargs["max_tokens"] = max_tokens
|
||||
|
||||
if stream:
|
||||
async def _stream_gen():
|
||||
# chat.completions.create() 返回 AsyncStream,不是 coroutine
|
||||
# 不能 await,直接在 async for 中迭代
|
||||
async for chunk in self._client.chat.completions.create(**kwargs):
|
||||
delta = chunk.choices[0].delta.content
|
||||
if delta:
|
||||
yield delta
|
||||
return _stream_gen()
|
||||
|
||||
resp = await self._client.chat.completions.create(**kwargs)
|
||||
return resp.choices[0].message.content or ""
|
||||
|
||||
|
||||
def make_backend(name: str, cfg: BackendConfig) -> BaseBackend:
|
||||
"""工厂函数:根据 type 创建对应后端实例。"""
|
||||
if cfg.type == "anthropic":
|
||||
return AnthropicBackend(cfg)
|
||||
if cfg.type in ("openai", "compatible"):
|
||||
return OpenAIBackend(cfg)
|
||||
if cfg.type == "ollama":
|
||||
# Ollama 提供 OpenAI 兼容 API,复用 OpenAIBackend
|
||||
cfg.base_url = cfg.base_url or "http://localhost:11434/v1"
|
||||
return OpenAIBackend(cfg)
|
||||
raise ValueError(f"未知后端类型: {cfg.type!r}(backend={name!r})")
|
||||
@@ -1,48 +0,0 @@
|
||||
"""语义复杂度评分器:query 与复杂任务模板集的最大余弦相似度。"""
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
from pathlib import Path
|
||||
|
||||
import numpy as np
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
class ComplexityScorer:
|
||||
"""启动时预计算模板向量,运行时 <2ms 评分。"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
embedder,
|
||||
complex_templates: list[str],
|
||||
threshold: float = 0.72,
|
||||
) -> None:
|
||||
self._threshold = threshold
|
||||
if complex_templates:
|
||||
vecs = embedder.embed(complex_templates, prefix="passage")
|
||||
self._template_vecs: np.ndarray = vecs # (N, dim)
|
||||
else:
|
||||
self._template_vecs = np.empty((0, 1), dtype=np.float32)
|
||||
self._embedder = embedder
|
||||
|
||||
def score(self, query: str) -> float:
|
||||
"""返回 0-1 复杂度分数。"""
|
||||
if self._template_vecs.shape[0] == 0:
|
||||
return 0.0
|
||||
q_vec = self._embedder.embed_query(query) # (dim,)
|
||||
sims = self._template_vecs @ q_vec # (N,)
|
||||
return float(sims.max())
|
||||
|
||||
def is_complex(self, query: str) -> bool:
|
||||
return self.score(query) > self._threshold
|
||||
|
||||
@staticmethod
|
||||
def load_templates(path: Path) -> list[str]:
|
||||
"""从文件加载模板,忽略空行和 # 注释行。"""
|
||||
templates = []
|
||||
for line in path.read_text(encoding="utf-8").splitlines():
|
||||
line = line.strip()
|
||||
if line and not line.startswith("#"):
|
||||
templates.append(line)
|
||||
return templates
|
||||
@@ -1,157 +0,0 @@
|
||||
"""token 预算管理与上下文压缩器(五层优先级)。"""
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
from dataclasses import dataclass, field
|
||||
from typing import Any
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
# tiktoken 可选,回退到字符估算
|
||||
try:
|
||||
import tiktoken
|
||||
_ENC = tiktoken.get_encoding("cl100k_base")
|
||||
def count_tokens(text: str) -> int:
|
||||
return len(_ENC.encode(text))
|
||||
except Exception:
|
||||
def count_tokens(text: str) -> int: # type: ignore[misc]
|
||||
return max(1, len(text) // 4)
|
||||
|
||||
|
||||
@dataclass
|
||||
class CompressResult:
|
||||
facts: list[dict[str, Any]]
|
||||
long_term: list[dict[str, Any]]
|
||||
short_term: list[dict[str, Any]]
|
||||
formatted: str
|
||||
tokens_used: int
|
||||
budget: dict[str, int]
|
||||
|
||||
|
||||
class Compressor:
|
||||
"""五层优先级 token 预算压缩器。
|
||||
|
||||
优先级(从高到低):
|
||||
1. facts(事实) — 语义检索,category 独立 top-k
|
||||
2. working_memory — 最近 N 轮,无条件注入
|
||||
3. summary — 当前 session 摘要(在 short_term 中 role='summary')
|
||||
4. history — 语义历史(short_term 中非 summary 且非 working_memory)
|
||||
5. knowledge(long_term)— 文档知识库
|
||||
"""
|
||||
|
||||
FACTS_RATIO = 0.25
|
||||
LONG_RATIO = 0.25
|
||||
SHORT_RATIO = 0.35
|
||||
|
||||
def __init__(self, token_budget: int = 2000) -> None:
|
||||
self._budget = token_budget
|
||||
|
||||
def compress(
|
||||
self,
|
||||
facts: list[dict[str, Any]],
|
||||
working_memory: list[dict[str, Any]],
|
||||
long_term: list[dict[str, Any]],
|
||||
short_term: list[dict[str, Any]],
|
||||
) -> CompressResult:
|
||||
facts_budget = int(self._budget * self.FACTS_RATIO)
|
||||
long_budget = int(self._budget * self.LONG_RATIO)
|
||||
short_budget = int(self._budget * self.SHORT_RATIO)
|
||||
|
||||
# facts:按预算截断
|
||||
selected_facts = self._trim_items(facts, "content", facts_budget)
|
||||
|
||||
# working_memory:计算 token 消耗(无条件注入)
|
||||
wm_tokens = sum(count_tokens(t.get("content", "")) for t in working_memory)
|
||||
|
||||
# working_memory 中的 id 集合,用于去重
|
||||
wm_ids = {t.get("id") for t in working_memory if t.get("id") is not None}
|
||||
|
||||
# summary 单独提取(置于 short_term 最前)
|
||||
summaries = [t for t in short_term if t.get("role") == "summary"]
|
||||
# history = short_term 中非 summary 且不在 working_memory 中
|
||||
history = [
|
||||
t for t in short_term
|
||||
if t.get("role") != "summary" and t.get("id") not in wm_ids
|
||||
]
|
||||
|
||||
remaining_short = max(0, short_budget - wm_tokens)
|
||||
selected_summaries = self._trim_items(summaries, "content", remaining_short // 2)
|
||||
selected_history = self._trim_items(
|
||||
history, "content",
|
||||
max(0, remaining_short - sum(count_tokens(t.get("content", "")) for t in selected_summaries))
|
||||
)
|
||||
|
||||
selected_long = self._trim_items(long_term, "chunk", long_budget)
|
||||
|
||||
formatted = self._format(
|
||||
selected_facts, working_memory, selected_summaries, selected_history, selected_long
|
||||
)
|
||||
tokens_used = count_tokens(formatted)
|
||||
|
||||
return CompressResult(
|
||||
facts=selected_facts,
|
||||
long_term=selected_long,
|
||||
short_term=selected_summaries + selected_history,
|
||||
formatted=formatted,
|
||||
tokens_used=tokens_used,
|
||||
budget={
|
||||
"facts": facts_budget,
|
||||
"long_term": long_budget,
|
||||
"short_term": short_budget,
|
||||
"used": tokens_used,
|
||||
},
|
||||
)
|
||||
|
||||
def _trim_items(
|
||||
self, items: list[dict], text_key: str, budget: int
|
||||
) -> list[dict]:
|
||||
result, used = [], 0
|
||||
for item in items:
|
||||
n = count_tokens(item.get(text_key, ""))
|
||||
if used + n > budget:
|
||||
break
|
||||
result.append(item)
|
||||
used += n
|
||||
return result
|
||||
|
||||
def _format(
|
||||
self,
|
||||
facts: list[dict],
|
||||
working_memory: list[dict],
|
||||
summaries: list[dict],
|
||||
history: list[dict],
|
||||
long_term: list[dict],
|
||||
) -> str:
|
||||
parts: list[str] = []
|
||||
|
||||
if facts:
|
||||
lines = []
|
||||
for f in facts:
|
||||
cat = f.get("category", "other").upper()
|
||||
content = f.get("content", "")
|
||||
lines.append(f"{cat}: {content}")
|
||||
parts.append(f"<facts>\n{chr(10).join(lines)}\n</facts>")
|
||||
|
||||
if working_memory:
|
||||
wm_lines = "\n".join(
|
||||
f"{t.get('role', 'unknown').upper()}: {t.get('content', '')}"
|
||||
for t in working_memory
|
||||
)
|
||||
parts.append(f"<working_memory>\n{wm_lines}\n</working_memory>")
|
||||
|
||||
short_items = summaries + history
|
||||
if short_items:
|
||||
hist_lines = "\n".join(
|
||||
f"{t.get('role', 'unknown').upper()}: {t.get('content', '')}"
|
||||
for t in short_items
|
||||
)
|
||||
parts.append(f"<history>\n{hist_lines}\n</history>")
|
||||
|
||||
if long_term:
|
||||
chunks = "\n\n".join(
|
||||
f"[{item.get('path', '?')}]\n{item.get('chunk', '')}"
|
||||
for item in long_term
|
||||
)
|
||||
parts.append(f"<knowledge>\n{chunks}\n</knowledge>")
|
||||
|
||||
return "\n\n".join(parts)
|
||||
@@ -1,157 +0,0 @@
|
||||
"""mem-bridge 配置:YAML 驱动的 dataclass。"""
|
||||
from __future__ import annotations
|
||||
|
||||
import os
|
||||
from dataclasses import dataclass, field
|
||||
from pathlib import Path
|
||||
from typing import Any
|
||||
|
||||
import yaml
|
||||
|
||||
|
||||
@dataclass
|
||||
class BackendConfig:
|
||||
type: str # anthropic | openai | compatible
|
||||
model: str
|
||||
api_key: str = ""
|
||||
base_url: str = "" # compatible 类型使用
|
||||
|
||||
|
||||
@dataclass
|
||||
class BridgeConfig:
|
||||
# memory-service
|
||||
memory_host: str = "0.0.0.0"
|
||||
memory_port: int = 8001
|
||||
db_dir: Path = field(default_factory=lambda: Path("~/.embed_db").expanduser())
|
||||
short_term_top_k: int = 5
|
||||
long_term_top_k: int = 3
|
||||
recent_turns_verbatim: int = 2
|
||||
|
||||
# router-service
|
||||
router_host: str = "0.0.0.0"
|
||||
router_port: int = 8000
|
||||
token_budget: int = 2000
|
||||
complexity_threshold: float = 0.72
|
||||
complexity_templates_path: Path = field(
|
||||
default_factory=lambda: Path("config/complexity_templates.txt")
|
||||
)
|
||||
default_backend: str = ""
|
||||
heavy_backend: str = ""
|
||||
backends: dict[str, BackendConfig] = field(default_factory=dict)
|
||||
fallback_chain: list[str] = field(default_factory=list)
|
||||
fallback_timeout: float = 10.0
|
||||
routing_rules: dict[str, str] = field(default_factory=dict)
|
||||
|
||||
# 摘要归档配置
|
||||
summarization_enabled: bool = True
|
||||
summarization_window_size: int = 20
|
||||
summarization_backend: str = "" # 空时使用 default_backend
|
||||
|
||||
# Query 改写配置
|
||||
query_rewrite_enabled: bool = False
|
||||
query_rewrite_backend: str = "" # 空时使用 default_backend
|
||||
|
||||
# System Prompt 模板配置
|
||||
prompt_templates_dir: Path = field(
|
||||
default_factory=lambda: Path("config/prompts")
|
||||
)
|
||||
prompt_templates_active: list[str] = field(default_factory=list)
|
||||
|
||||
# 事实提取配置
|
||||
fact_extraction_enabled: bool = False
|
||||
fact_extraction_backend: str = ""
|
||||
fact_extraction_window_turns: int = 6
|
||||
fact_extraction_dedup_threshold: float = 0.92
|
||||
|
||||
# 混合检索配置
|
||||
hybrid_search_enabled: bool = True
|
||||
hybrid_search_bm25_weight: float = 0.3
|
||||
|
||||
# 时间衰减配置
|
||||
time_decay_enabled: bool = True
|
||||
time_decay_lambda: float = 0.05
|
||||
|
||||
# Working Memory 层
|
||||
working_memory_turns: int = 2
|
||||
|
||||
# facts category 独立 top-k
|
||||
fact_category_top_k: dict[str, int] = field(
|
||||
default_factory=lambda: {
|
||||
"preference": 5,
|
||||
"background": 5,
|
||||
"goal": 3,
|
||||
"habit": 3,
|
||||
"other": 2,
|
||||
}
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def from_yaml(cls, path: Path) -> BridgeConfig:
|
||||
raw = yaml.safe_load(path.read_text(encoding="utf-8"))
|
||||
mem = raw.get("memory_service", {})
|
||||
rtr = raw.get("router_service", {})
|
||||
complexity = rtr.get("complexity", {})
|
||||
|
||||
backends: dict[str, BackendConfig] = {}
|
||||
for name, bcfg in rtr.get("backends", {}).items():
|
||||
if name in ("default", "heavy"):
|
||||
continue
|
||||
api_key = bcfg.get("api_key", "")
|
||||
# 展开环境变量 ${VAR}
|
||||
if api_key.startswith("${") and api_key.endswith("}"):
|
||||
api_key = os.environ.get(api_key[2:-1], "")
|
||||
backends[name] = BackendConfig(
|
||||
type=bcfg["type"],
|
||||
model=bcfg["model"],
|
||||
api_key=api_key,
|
||||
base_url=bcfg.get("base_url", ""),
|
||||
)
|
||||
|
||||
summ = mem.get("summarization", {})
|
||||
qr = rtr.get("query_rewrite", {})
|
||||
pt = rtr.get("prompt_templates", {})
|
||||
|
||||
_fe = mem.get("fact_extraction", {})
|
||||
_hs = mem.get("hybrid_search", {})
|
||||
_td = mem.get("time_decay", {})
|
||||
_default_cat_top_k = {"preference": 5, "background": 5, "goal": 3, "habit": 3, "other": 2}
|
||||
|
||||
db_dir_raw = mem.get("db_dir", "~/.embed_db")
|
||||
return cls(
|
||||
memory_host=mem.get("host", "0.0.0.0"),
|
||||
memory_port=mem.get("port", 8001),
|
||||
db_dir=Path(db_dir_raw).expanduser(),
|
||||
short_term_top_k=mem.get("short_term_top_k", 5),
|
||||
long_term_top_k=mem.get("long_term_top_k", 3),
|
||||
recent_turns_verbatim=mem.get("recent_turns_verbatim", 2),
|
||||
router_host=rtr.get("host", "0.0.0.0"),
|
||||
router_port=rtr.get("port", 8000),
|
||||
token_budget=rtr.get("token_budget", 2000),
|
||||
complexity_threshold=complexity.get("threshold", 0.72),
|
||||
complexity_templates_path=Path(
|
||||
complexity.get("templates", "config/complexity_templates.txt")
|
||||
),
|
||||
default_backend=rtr.get("backends", {}).get("default", ""),
|
||||
heavy_backend=rtr.get("backends", {}).get("heavy", ""),
|
||||
backends=backends,
|
||||
fallback_chain=rtr.get("fallback_chain", []),
|
||||
fallback_timeout=rtr.get("fallback_timeout", 10.0),
|
||||
routing_rules=rtr.get("routing_rules", {}),
|
||||
summarization_enabled=summ.get("enabled", True),
|
||||
summarization_window_size=summ.get("window_size", 20),
|
||||
summarization_backend=summ.get("backend") or "",
|
||||
query_rewrite_enabled=qr.get("enabled", False),
|
||||
query_rewrite_backend=qr.get("backend") or "",
|
||||
prompt_templates_dir=Path(pt.get("dir") or "config/prompts"),
|
||||
prompt_templates_active=pt.get("active") or [],
|
||||
fact_extraction_enabled=_fe.get("enabled", False),
|
||||
fact_extraction_backend=_fe.get("backend") or "",
|
||||
fact_extraction_window_turns=_fe.get("window_turns", 6),
|
||||
fact_extraction_dedup_threshold=_fe.get("dedup_threshold", 0.92),
|
||||
hybrid_search_enabled=_hs.get("enabled", True),
|
||||
hybrid_search_bm25_weight=_hs.get("bm25_weight", 0.3),
|
||||
time_decay_enabled=_td.get("enabled", True),
|
||||
time_decay_lambda=_td.get("lambda", 0.05),
|
||||
working_memory_turns=mem.get("working_memory_turns", 2),
|
||||
fact_category_top_k={**_default_cat_top_k, **(mem.get("fact_category_top_k") or {})},
|
||||
)
|
||||
@@ -1,71 +0,0 @@
|
||||
"""会话上下文管理:滑动窗口 + 异步摘要归档。"""
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
from typing import TYPE_CHECKING
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from mem_bridge.backends import BaseBackend
|
||||
from mem_bridge.turn_store import TurnStore
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
_SUMMARY_SYSTEM = "你是对话摘要助手。"
|
||||
_SUMMARY_PROMPT = (
|
||||
"请将以下对话历史压缩为1-2段摘要,保留关键信息、决定和结论,省略闲聊:\n\n{turns}"
|
||||
)
|
||||
|
||||
|
||||
class ContextManager:
|
||||
"""滑动窗口 + 异步摘要归档。
|
||||
|
||||
当会话普通轮次超出 window_size 时,将旧轮次摘要后归档为
|
||||
role='summary' 的特殊轮次,并删除原始旧轮次。
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
turn_store: TurnStore,
|
||||
backend: BaseBackend,
|
||||
window_size: int = 20,
|
||||
) -> None:
|
||||
self._store = turn_store
|
||||
self._backend = backend
|
||||
self._window_size = window_size
|
||||
|
||||
async def check_and_summarize(self, session_id: str) -> None:
|
||||
"""检查会话是否超出滑动窗口,超出则生成摘要归档旧轮次。"""
|
||||
turns = self._store.list_turns(session_id)
|
||||
regular = [t for t in turns if t["role"] != "summary"]
|
||||
if len(regular) <= self._window_size:
|
||||
return
|
||||
old_turns = regular[: len(regular) - self._window_size]
|
||||
await self._summarize_and_archive(session_id, old_turns)
|
||||
|
||||
async def _summarize_and_archive(
|
||||
self, session_id: str, old_turns: list[dict]
|
||||
) -> None:
|
||||
"""调用 LLM 生成摘要,存入 TurnStore,删除旧轮次。"""
|
||||
from mem_bridge.models import ChatMessage # 延迟加载,避免循环 import
|
||||
|
||||
turns_text = "\n".join(
|
||||
f"{t['role'].upper()}: {t['content']}" for t in old_turns
|
||||
)
|
||||
messages = [
|
||||
ChatMessage(role="system", content=_SUMMARY_SYSTEM),
|
||||
ChatMessage(
|
||||
role="user",
|
||||
content=_SUMMARY_PROMPT.format(turns=turns_text),
|
||||
),
|
||||
]
|
||||
try:
|
||||
summary = str(await self._backend.chat(messages, stream=False))
|
||||
self._store.add_turn(session_id, "summary", summary)
|
||||
ids = [t["id"] for t in old_turns]
|
||||
self._store.delete_turns_by_ids(ids)
|
||||
logger.info(
|
||||
"session %s: 归档 %d 轮,生成摘要(%d 字)",
|
||||
session_id, len(old_turns), len(str(summary)),
|
||||
)
|
||||
except Exception as e:
|
||||
logger.warning("摘要生成失败,跳过归档: %s", e)
|
||||
@@ -1,129 +0,0 @@
|
||||
"""LLM 异步提取对话事实,写入 FactStore。"""
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
import logging
|
||||
import re
|
||||
from typing import TYPE_CHECKING, Any
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from mem_bridge.backends import BaseBackend
|
||||
from mem_bridge.fact_store import FactStore
|
||||
|
||||
import numpy as np
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
_EXTRACT_SYSTEM = (
|
||||
"你是事实提取助手。从对话中提取用户的持久化事实。"
|
||||
"只提取明确信息,不推断,不重复已知事实。"
|
||||
"category 枚举: preference/background/goal/habit/other。"
|
||||
"输出 JSON 数组,不要其他文字。"
|
||||
)
|
||||
_EXTRACT_PROMPT = (
|
||||
"从以下对话中提取用户的持久化事实(姓名、偏好、背景、目标、习惯等)。\n"
|
||||
"若无可提取事实,返回空数组 []。\n"
|
||||
"格式:[{{\"content\":\"...\", \"category\":\"...\", \"entity\":\"user\","
|
||||
" \"tags\":[\"...\"], \"confidence\":0.9}}]\n\n对话:\n{turns_text}"
|
||||
)
|
||||
|
||||
|
||||
class FactExtractor:
|
||||
"""LLM 异步提取对话事实,写入 FactStore。"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
backend: BaseBackend,
|
||||
fact_store: FactStore,
|
||||
embedder: Any,
|
||||
write_lock: Any,
|
||||
enabled: bool = False,
|
||||
window_turns: int = 6,
|
||||
dedup_threshold: float = 0.92,
|
||||
) -> None:
|
||||
self._backend = backend
|
||||
self._fact_store = fact_store
|
||||
self._embedder = embedder
|
||||
self._write_lock = write_lock
|
||||
self._enabled = enabled
|
||||
self._window_turns = window_turns
|
||||
self._dedup_threshold = dedup_threshold
|
||||
|
||||
async def extract_and_store(self, session_id: str, turns: list[dict]) -> int:
|
||||
"""从最近轮次提取事实,返回新增条数。失败静默跳过。"""
|
||||
if not self._enabled:
|
||||
return 0
|
||||
recent = turns[-self._window_turns:]
|
||||
if not recent:
|
||||
return 0
|
||||
|
||||
turns_text = "\n".join(
|
||||
f"{t['role'].upper()}: {t['content']}" for t in recent
|
||||
)
|
||||
from mem_bridge.models import ChatMessage
|
||||
|
||||
messages = [
|
||||
ChatMessage(role="system", content=_EXTRACT_SYSTEM),
|
||||
ChatMessage(role="user", content=_EXTRACT_PROMPT.format(turns_text=turns_text)),
|
||||
]
|
||||
try:
|
||||
raw = str(await self._backend.chat(messages, stream=False))
|
||||
candidates = self._parse_json(raw)
|
||||
return await self._store_candidates(session_id, candidates)
|
||||
except Exception as e:
|
||||
logger.warning("事实提取失败,跳过: %s", e)
|
||||
return 0
|
||||
|
||||
def _parse_json(self, raw: str) -> list[dict]:
|
||||
"""容错解析 LLM 返回的 JSON 数组。"""
|
||||
match = re.search(r"\[.*\]", raw, re.DOTALL)
|
||||
if not match:
|
||||
return []
|
||||
try:
|
||||
result = json.loads(match.group())
|
||||
if isinstance(result, list):
|
||||
return result
|
||||
except json.JSONDecodeError:
|
||||
pass
|
||||
return []
|
||||
|
||||
async def _store_candidates(
|
||||
self, session_id: str, candidates: list[dict]
|
||||
) -> int:
|
||||
added = 0
|
||||
for c in candidates:
|
||||
content = str(c.get("content", "")).strip()
|
||||
if not content:
|
||||
continue
|
||||
category = str(c.get("category", "other"))
|
||||
entity = c.get("entity") or None
|
||||
tags = c.get("tags", [])
|
||||
if not isinstance(tags, list):
|
||||
tags = []
|
||||
confidence = float(c.get("confidence", 1.0))
|
||||
|
||||
vec = np.array(self._embedder.embed_query(content), dtype=np.float32)
|
||||
|
||||
# 语义去重检查 + 写入全部纳入锁保护,避免并发 TOCTOU 竞态
|
||||
async with self._write_lock:
|
||||
ids, dists = self._fact_store.faiss_search(vec, top_k=1)
|
||||
if ids.size > 0 and float(dists[0]) >= self._dedup_threshold:
|
||||
existing = self._fact_store.search_by_faiss_ids([int(ids[0])])
|
||||
if existing:
|
||||
old = existing[0]
|
||||
if len(content) > len(old["content"]):
|
||||
# 新内容更具体 → 更新
|
||||
self._fact_store.update_fact(old["id"], content, confidence)
|
||||
logger.debug("更新事实 id=%d: %r", old["id"], content[:50])
|
||||
elif content != old["content"]:
|
||||
# 相似但不更具体 → 标记冲突
|
||||
self._fact_store.mark_conflict(old["id"])
|
||||
logger.debug("标记冲突 id=%d", old["id"])
|
||||
continue
|
||||
self._fact_store.add_fact(
|
||||
content, category, entity, tags,
|
||||
session_id, vec, confidence,
|
||||
)
|
||||
added += 1
|
||||
logger.debug("新增事实: category=%s content=%r", category, content[:50])
|
||||
return added
|
||||
@@ -1,230 +0,0 @@
|
||||
"""跨 session 持久化事实存储:SQLite + facts.faiss。"""
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
import logging
|
||||
import sqlite3
|
||||
import time
|
||||
from pathlib import Path
|
||||
|
||||
import numpy as np
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
_DDL = """
|
||||
CREATE TABLE IF NOT EXISTS facts (
|
||||
id INTEGER PRIMARY KEY AUTOINCREMENT,
|
||||
content TEXT NOT NULL,
|
||||
category TEXT NOT NULL DEFAULT 'other',
|
||||
entity TEXT,
|
||||
tags TEXT DEFAULT '[]',
|
||||
source_session TEXT NOT NULL,
|
||||
faiss_id INTEGER UNIQUE,
|
||||
confidence REAL DEFAULT 1.0,
|
||||
conflict INTEGER DEFAULT 0,
|
||||
created_at REAL NOT NULL,
|
||||
updated_at REAL NOT NULL
|
||||
);
|
||||
CREATE INDEX IF NOT EXISTS idx_facts_category ON facts(category);
|
||||
CREATE INDEX IF NOT EXISTS idx_facts_entity ON facts(entity);
|
||||
CREATE INDEX IF NOT EXISTS idx_facts_faiss ON facts(faiss_id);
|
||||
CREATE VIRTUAL TABLE IF NOT EXISTS facts_fts USING fts5(
|
||||
content, category, entity, tags,
|
||||
content='facts', content_rowid='id'
|
||||
);
|
||||
CREATE TRIGGER IF NOT EXISTS facts_ai AFTER INSERT ON facts BEGIN
|
||||
INSERT INTO facts_fts(rowid, content, category, entity, tags)
|
||||
VALUES (new.id, new.content, new.category, new.entity, new.tags);
|
||||
END;
|
||||
CREATE TRIGGER IF NOT EXISTS facts_ad AFTER DELETE ON facts BEGIN
|
||||
INSERT INTO facts_fts(facts_fts, rowid, content, category, entity, tags)
|
||||
VALUES ('delete', old.id, old.content, old.category, old.entity, old.tags);
|
||||
END;
|
||||
CREATE TRIGGER IF NOT EXISTS facts_au AFTER UPDATE ON facts BEGIN
|
||||
INSERT INTO facts_fts(facts_fts, rowid, content, category, entity, tags)
|
||||
VALUES ('delete', old.id, old.content, old.category, old.entity, old.tags);
|
||||
INSERT INTO facts_fts(rowid, content, category, entity, tags)
|
||||
VALUES (new.id, new.content, new.category, new.entity, new.tags);
|
||||
END;
|
||||
"""
|
||||
|
||||
|
||||
class FactStore:
|
||||
"""跨 session 持久化事实存储,SQLite + facts.faiss(IndexFlatIP)。"""
|
||||
|
||||
def __init__(self, db_path: Path, index_path: Path, emb_dim: int = 384) -> None:
|
||||
self._db_path = db_path
|
||||
self._index_path = index_path
|
||||
self._emb_dim = emb_dim
|
||||
self._index = None
|
||||
self._init()
|
||||
|
||||
def _connect(self) -> sqlite3.Connection:
|
||||
conn = sqlite3.connect(self._db_path)
|
||||
conn.row_factory = sqlite3.Row
|
||||
conn.execute("PRAGMA journal_mode=WAL;")
|
||||
conn.execute("PRAGMA synchronous=NORMAL;")
|
||||
return conn
|
||||
|
||||
def _init(self) -> None:
|
||||
with self._connect() as conn:
|
||||
conn.executescript(_DDL)
|
||||
|
||||
def load_index(self) -> None:
|
||||
"""启动时从磁盘加载 facts.faiss(mmap=False 保持可写)。"""
|
||||
import faiss # type: ignore[import]
|
||||
|
||||
if self._index_path.exists():
|
||||
try:
|
||||
self._index = faiss.read_index(str(self._index_path))
|
||||
logger.info("facts.faiss 已加载: ntotal=%d", self._index.ntotal)
|
||||
except Exception as e:
|
||||
logger.warning("facts.faiss 加载失败,从空索引开始: %s", e)
|
||||
|
||||
def _ensure_index(self) -> None:
|
||||
if self._index is None:
|
||||
import faiss # type: ignore[import]
|
||||
flat = faiss.IndexFlatIP(self._emb_dim)
|
||||
self._index = faiss.IndexIDMap2(flat)
|
||||
|
||||
def _save_index(self) -> None:
|
||||
if self._index is None:
|
||||
return
|
||||
import faiss # type: ignore[import]
|
||||
self._index_path.parent.mkdir(parents=True, exist_ok=True)
|
||||
faiss.write_index(self._index, str(self._index_path))
|
||||
|
||||
def next_faiss_id(self) -> int:
|
||||
with self._connect() as conn:
|
||||
row = conn.execute("SELECT MAX(faiss_id) FROM facts").fetchone()
|
||||
return (row[0] + 1) if row[0] is not None else 0
|
||||
|
||||
def add_fact(
|
||||
self,
|
||||
content: str,
|
||||
category: str,
|
||||
entity: str | None,
|
||||
tags: list[str],
|
||||
source_session: str,
|
||||
vec: np.ndarray,
|
||||
confidence: float = 1.0,
|
||||
conflict: bool = False,
|
||||
) -> int:
|
||||
"""写入事实,返回 fact id。调用方需在写锁内调用此方法。"""
|
||||
self._ensure_index()
|
||||
faiss_id = self.next_faiss_id()
|
||||
now = time.time()
|
||||
vec_arr = vec.astype(np.float32).reshape(1, -1)
|
||||
id_arr = np.array([faiss_id], dtype=np.int64)
|
||||
self._index.add_with_ids(vec_arr, id_arr) # type: ignore[union-attr]
|
||||
self._save_index()
|
||||
with self._connect() as conn:
|
||||
cur = conn.execute(
|
||||
"INSERT INTO facts"
|
||||
"(content,category,entity,tags,source_session,"
|
||||
"faiss_id,confidence,conflict,created_at,updated_at)"
|
||||
" VALUES (?,?,?,?,?,?,?,?,?,?)",
|
||||
(
|
||||
content, category, entity,
|
||||
json.dumps(tags, ensure_ascii=False),
|
||||
source_session, faiss_id, confidence,
|
||||
int(conflict), now, now,
|
||||
),
|
||||
)
|
||||
return cur.lastrowid # type: ignore[return-value]
|
||||
|
||||
def update_fact(self, fact_id: int, content: str, confidence: float) -> None:
|
||||
"""更新事实内容(去重时发现新内容更具体时调用)。"""
|
||||
now = time.time()
|
||||
with self._connect() as conn:
|
||||
conn.execute(
|
||||
"UPDATE facts SET content=?, confidence=?, updated_at=? WHERE id=?",
|
||||
(content, confidence, now, fact_id),
|
||||
)
|
||||
|
||||
def mark_conflict(self, fact_id: int) -> None:
|
||||
"""标记事实存在冲突(保留双方,上层注入时取最新)。"""
|
||||
now = time.time()
|
||||
with self._connect() as conn:
|
||||
conn.execute(
|
||||
"UPDATE facts SET conflict=1, updated_at=? WHERE id=?",
|
||||
(now, fact_id),
|
||||
)
|
||||
|
||||
def search_by_faiss_ids(self, faiss_ids: list[int]) -> list[dict]:
|
||||
if not faiss_ids:
|
||||
return []
|
||||
placeholders = ",".join("?" * len(faiss_ids))
|
||||
with self._connect() as conn:
|
||||
rows = conn.execute(
|
||||
f"SELECT * FROM facts WHERE faiss_id IN ({placeholders})",
|
||||
faiss_ids,
|
||||
).fetchall()
|
||||
return [dict(r) for r in rows]
|
||||
|
||||
def faiss_search(
|
||||
self, query_vec: np.ndarray, top_k: int = 20
|
||||
) -> tuple[np.ndarray, np.ndarray]:
|
||||
"""语义搜索,返回 (ids, distances),空索引时返回空数组。"""
|
||||
if self._index is None or self._index.ntotal == 0:
|
||||
return np.array([], dtype=np.int64), np.array([], dtype=np.float32)
|
||||
query = query_vec.astype(np.float32).reshape(1, -1)
|
||||
k = min(top_k, self._index.ntotal)
|
||||
distances, ids = self._index.search(query, k)
|
||||
valid = ids[0] >= 0
|
||||
return ids[0][valid], distances[0][valid]
|
||||
|
||||
def bm25_search(self, query: str, limit: int = 20) -> list[dict]:
|
||||
"""FTS5 BM25 全文检索(bm25_score 为负值,越负越相关)。"""
|
||||
tokens = [t for t in query.split() if t]
|
||||
if not tokens:
|
||||
return []
|
||||
fts_query = " OR ".join(tokens)
|
||||
with self._connect() as conn:
|
||||
rows = conn.execute(
|
||||
"""
|
||||
SELECT f.id, f.faiss_id, f.content, f.category, f.entity,
|
||||
f.tags, f.confidence, f.conflict,
|
||||
f.created_at, f.updated_at,
|
||||
bm25(facts_fts) AS bm25_score
|
||||
FROM facts_fts
|
||||
JOIN facts f ON f.id = facts_fts.rowid
|
||||
WHERE facts_fts MATCH ?
|
||||
ORDER BY bm25_score
|
||||
LIMIT ?
|
||||
""",
|
||||
(fts_query, limit),
|
||||
).fetchall()
|
||||
return [dict(r) for r in rows]
|
||||
|
||||
def list_facts(
|
||||
self,
|
||||
category: str | None = None,
|
||||
entity: str | None = None,
|
||||
) -> list[dict]:
|
||||
conditions: list[str] = []
|
||||
params: list[str] = []
|
||||
if category:
|
||||
conditions.append("category=?")
|
||||
params.append(category)
|
||||
if entity:
|
||||
conditions.append("entity=?")
|
||||
params.append(entity)
|
||||
where = f"WHERE {' AND '.join(conditions)}" if conditions else ""
|
||||
with self._connect() as conn:
|
||||
rows = conn.execute(
|
||||
f"SELECT * FROM facts {where} ORDER BY updated_at DESC",
|
||||
params,
|
||||
).fetchall()
|
||||
return [dict(r) for r in rows]
|
||||
|
||||
def delete_fact(self, fact_id: int) -> None:
|
||||
with self._connect() as conn:
|
||||
row = conn.execute(
|
||||
"SELECT faiss_id FROM facts WHERE id=?", (fact_id,)
|
||||
).fetchone()
|
||||
conn.execute("DELETE FROM facts WHERE id=?", (fact_id,))
|
||||
if row and row["faiss_id"] is not None and self._index is not None:
|
||||
ids_to_remove = np.array([row["faiss_id"]], dtype=np.int64)
|
||||
self._index.remove_ids(ids_to_remove)
|
||||
self._save_index()
|
||||
@@ -1,520 +0,0 @@
|
||||
"""Memory Service FastAPI 应用:对话轮次存取 + 上下文压缩检索。"""
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import logging
|
||||
import math
|
||||
from pathlib import Path
|
||||
from typing import TYPE_CHECKING, Any, Coroutine
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from mem_bridge.context_manager import ContextManager
|
||||
from mem_bridge.fact_extractor import FactExtractor
|
||||
|
||||
import numpy as np
|
||||
from fastapi import FastAPI
|
||||
from pydantic import BaseModel
|
||||
|
||||
from embed_db.config import Config
|
||||
from embed_db.index import VectorIndex
|
||||
from embed_db.pipeline import Pipeline
|
||||
from mem_bridge.compressor import Compressor
|
||||
from mem_bridge.config import BridgeConfig
|
||||
from mem_bridge.fact_store import FactStore
|
||||
from mem_bridge.turn_store import TurnStore
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
# 模块级全局 bridge_cfg,供单元测试直接操作(由 create_memory_app 注入)
|
||||
bridge_cfg: BridgeConfig = BridgeConfig()
|
||||
|
||||
|
||||
def _log_task_exception(task: asyncio.Task) -> None: # type: ignore[type-arg]
|
||||
"""后台 task 完成回调:记录异常,避免异常静默丢失。"""
|
||||
if task.cancelled():
|
||||
return
|
||||
exc = task.exception()
|
||||
if exc is not None:
|
||||
logger.warning("后台 task 异常: %s", exc, exc_info=exc)
|
||||
|
||||
|
||||
def _safe_create_task(coro: Coroutine[Any, Any, Any]) -> None:
|
||||
"""在已运行的事件循环中创建 task,无循环时静默丢弃。"""
|
||||
try:
|
||||
loop = asyncio.get_running_loop()
|
||||
task = loop.create_task(coro)
|
||||
task.add_done_callback(_log_task_exception)
|
||||
except RuntimeError:
|
||||
try:
|
||||
coro.close()
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
|
||||
def _hybrid_scores(
|
||||
semantic_ids: np.ndarray,
|
||||
semantic_dists: np.ndarray,
|
||||
bm25_rows: list[dict],
|
||||
bm25_weight: float,
|
||||
) -> dict[int, float]:
|
||||
"""将语义分数和 BM25 分数归一化后混合,返回 {faiss_id: score}。"""
|
||||
scores: dict[int, float] = {}
|
||||
|
||||
# 语义分数(已归一化到 [0,1],直接使用)
|
||||
for fid, dist in zip(semantic_ids, semantic_dists):
|
||||
scores[int(fid)] = float(dist) * (1 - bm25_weight)
|
||||
|
||||
# BM25 分数(负值,越负越相关,归一化到 [0,1])
|
||||
if bm25_rows:
|
||||
min_raw = min(r["bm25_score"] for r in bm25_rows)
|
||||
max_raw = max(r["bm25_score"] for r in bm25_rows)
|
||||
span = max_raw - min_raw
|
||||
for r in bm25_rows:
|
||||
fid = r.get("faiss_id")
|
||||
if fid is None:
|
||||
continue
|
||||
if span == 0.0:
|
||||
# 单条结果,赋予中性权重
|
||||
norm_bm25 = 0.5
|
||||
else:
|
||||
norm = (r["bm25_score"] - min_raw) / span
|
||||
norm_bm25 = 1.0 - norm # 越负越相关,翻转为越大越好
|
||||
scores[int(fid)] = scores.get(int(fid), 0.0) + norm_bm25 * bm25_weight
|
||||
|
||||
return scores
|
||||
|
||||
|
||||
def _apply_time_decay(turns: list[dict], cfg: BridgeConfig | None = None) -> list[dict]:
|
||||
"""对历史轮次的 score 应用时间衰减。
|
||||
|
||||
Args:
|
||||
turns: 包含 score 和 created_at 字段的轮次列表。
|
||||
cfg: BridgeConfig 实例;为 None 时使用模块级全局 bridge_cfg。
|
||||
"""
|
||||
import time as _time
|
||||
_cfg = cfg if cfg is not None else bridge_cfg
|
||||
if not _cfg.time_decay_enabled:
|
||||
return turns
|
||||
now = _time.time()
|
||||
lam = _cfg.time_decay_lambda
|
||||
result = []
|
||||
for t in turns:
|
||||
age_hours = (now - t.get("created_at", now)) / 3600.0
|
||||
decay = math.exp(-lam * age_hours)
|
||||
result.append({**t, "score": float(t.get("score", 1.0)) * decay})
|
||||
return sorted(result, key=lambda x: x.get("score", 0.0), reverse=True)
|
||||
|
||||
|
||||
class AddTurnRequest(BaseModel):
|
||||
session_id: str
|
||||
role: str
|
||||
content: str
|
||||
|
||||
|
||||
class AddFactRequest(BaseModel):
|
||||
content: str
|
||||
category: str = "other"
|
||||
entity: str | None = None
|
||||
tags: list[str] = []
|
||||
|
||||
|
||||
def create_memory_app(
|
||||
bridge_cfg: BridgeConfig,
|
||||
embed_cfg: Config,
|
||||
embedder: Any | None = None,
|
||||
) -> FastAPI:
|
||||
import sys as _sys
|
||||
_sys.modules[__name__].bridge_cfg = bridge_cfg # 同步到模块级全局,便于单元测试访问
|
||||
app = FastAPI(title="mem-bridge memory-service")
|
||||
|
||||
# ── 基础组件 ──────────────────────────────────────
|
||||
bridge_cfg.db_dir.mkdir(parents=True, exist_ok=True)
|
||||
turn_store = TurnStore(bridge_cfg.db_dir / "turns.db")
|
||||
|
||||
turns_index_cfg = Config(
|
||||
db_dir=bridge_cfg.db_dir,
|
||||
model_path=embed_cfg.model_path,
|
||||
index_path=bridge_cfg.db_dir / "turns.faiss",
|
||||
meta_path=embed_cfg.meta_path,
|
||||
emb_dim=embed_cfg.emb_dim,
|
||||
nlist=embed_cfg.nlist,
|
||||
pq_m=embed_cfg.pq_m,
|
||||
min_vectors_for_ivf=embed_cfg.min_vectors_for_ivf,
|
||||
)
|
||||
turns_index = VectorIndex(turns_index_cfg)
|
||||
turns_index_path = bridge_cfg.db_dir / "turns.faiss"
|
||||
if turns_index_path.exists():
|
||||
try:
|
||||
turns_index.load(path=turns_index_path, mmap=False)
|
||||
except Exception as e:
|
||||
logger.warning("turns.faiss 加载失败,从空索引开始: %s", e)
|
||||
|
||||
# ── FactStore ──────────────────────────────────────
|
||||
fact_store = FactStore(
|
||||
db_path=bridge_cfg.db_dir / "facts.db",
|
||||
index_path=bridge_cfg.db_dir / "facts.faiss",
|
||||
emb_dim=embed_cfg.emb_dim,
|
||||
)
|
||||
fact_store.load_index()
|
||||
|
||||
# ── 写锁(懒初始化,Python 3.9 兼容)─────────────
|
||||
_write_lock: asyncio.Lock | None = None
|
||||
_facts_write_lock: asyncio.Lock | None = None
|
||||
|
||||
doc_pipeline = Pipeline(embed_cfg)
|
||||
if embedder is not None:
|
||||
doc_pipeline._embedder = embedder
|
||||
compressor = Compressor(token_budget=bridge_cfg.token_budget)
|
||||
|
||||
# ── ContextManager ─────────────────────────────────
|
||||
context_mgr: ContextManager | None = None
|
||||
if bridge_cfg.summarization_enabled and bridge_cfg.backends:
|
||||
summ_name = bridge_cfg.summarization_backend or bridge_cfg.default_backend
|
||||
if not summ_name:
|
||||
logger.warning("summarization_enabled=true 但未配置 default_backend,跳过 ContextManager")
|
||||
else:
|
||||
summ_cfg = bridge_cfg.backends.get(summ_name)
|
||||
if not summ_cfg:
|
||||
logger.warning("摘要后端 %r 未找到,跳过 ContextManager", summ_name)
|
||||
else:
|
||||
from mem_bridge.backends import make_backend
|
||||
from mem_bridge.context_manager import ContextManager
|
||||
summ_backend = make_backend(summ_name, summ_cfg)
|
||||
context_mgr = ContextManager(
|
||||
turn_store=turn_store,
|
||||
backend=summ_backend,
|
||||
window_size=bridge_cfg.summarization_window_size,
|
||||
)
|
||||
logger.info(
|
||||
"ContextManager 已初始化,后端=%r,窗口大小=%d",
|
||||
summ_name, bridge_cfg.summarization_window_size,
|
||||
)
|
||||
|
||||
# ── FactExtractor ──────────────────────────────────
|
||||
fact_extractor: FactExtractor | None = None
|
||||
if bridge_cfg.fact_extraction_enabled and bridge_cfg.backends:
|
||||
fe_name = bridge_cfg.fact_extraction_backend or bridge_cfg.default_backend
|
||||
if not fe_name:
|
||||
logger.warning("fact_extraction_enabled=true 但未配置 default_backend,跳过 FactExtractor")
|
||||
else:
|
||||
fe_cfg = bridge_cfg.backends.get(fe_name)
|
||||
if not fe_cfg:
|
||||
logger.warning("事实提取后端 %r 未找到,跳过 FactExtractor", fe_name)
|
||||
else:
|
||||
from mem_bridge.backends import make_backend
|
||||
from mem_bridge.fact_extractor import FactExtractor
|
||||
fe_backend = make_backend(fe_name, fe_cfg)
|
||||
|
||||
class _LazyFactExtractor:
|
||||
"""包装 FactExtractor,懒获取 _facts_write_lock。"""
|
||||
|
||||
def __init__(self) -> None:
|
||||
self._inner: FactExtractor | None = None
|
||||
self._backend = fe_backend
|
||||
|
||||
def _get_inner(self, lock: asyncio.Lock) -> FactExtractor:
|
||||
if self._inner is None:
|
||||
self._inner = FactExtractor(
|
||||
backend=self._backend,
|
||||
fact_store=fact_store,
|
||||
embedder=_get_embedder(),
|
||||
write_lock=lock,
|
||||
enabled=True,
|
||||
window_turns=bridge_cfg.fact_extraction_window_turns,
|
||||
dedup_threshold=bridge_cfg.fact_extraction_dedup_threshold,
|
||||
)
|
||||
return self._inner
|
||||
|
||||
async def extract_and_store(
|
||||
self, session_id: str, turns: list[dict], lock: asyncio.Lock
|
||||
) -> int:
|
||||
return await self._get_inner(lock).extract_and_store(
|
||||
session_id, turns
|
||||
)
|
||||
|
||||
fact_extractor = _LazyFactExtractor() # type: ignore[assignment]
|
||||
logger.info("FactExtractor 已配置,后端=%r", fe_name)
|
||||
|
||||
def _get_embedder() -> Any:
|
||||
if embedder is not None:
|
||||
return embedder
|
||||
doc_pipeline._ensure_embedder()
|
||||
return doc_pipeline._embedder
|
||||
|
||||
# ── 工具函数 ───────────────────────────────────────
|
||||
|
||||
def _hybrid_search_facts(q_vec: np.ndarray, query: str, top_k: int = 20) -> list[dict]:
|
||||
"""语义 + BM25 混合搜索事实,按 category top-k 筛选。"""
|
||||
sem_ids, sem_dists = fact_store.faiss_search(q_vec, top_k=top_k)
|
||||
bm25_rows: list[dict] = []
|
||||
if bridge_cfg.hybrid_search_enabled:
|
||||
try:
|
||||
bm25_rows = fact_store.bm25_search(query, limit=top_k)
|
||||
except Exception as e:
|
||||
logger.debug("BM25 搜索失败,退化为纯语义: %s", e)
|
||||
|
||||
scores = _hybrid_scores(
|
||||
sem_ids, sem_dists, bm25_rows,
|
||||
bridge_cfg.hybrid_search_bm25_weight if bridge_cfg.hybrid_search_enabled else 0.0,
|
||||
)
|
||||
|
||||
# 回查事实文本
|
||||
all_fids = list(scores.keys())
|
||||
facts = fact_store.search_by_faiss_ids(all_fids)
|
||||
for f in facts:
|
||||
f["score"] = scores.get(f.get("faiss_id", -1), 0.0)
|
||||
|
||||
# 按 category 独立 top-k
|
||||
cat_counts: dict[str, int] = {}
|
||||
selected: list[dict] = []
|
||||
cat_limits = bridge_cfg.fact_category_top_k
|
||||
default_limit = 2
|
||||
for f in sorted(facts, key=lambda x: x.get("score", 0.0), reverse=True):
|
||||
cat = f.get("category", "other")
|
||||
limit = cat_limits.get(cat, default_limit)
|
||||
if cat_counts.get(cat, 0) < limit:
|
||||
selected.append(f)
|
||||
cat_counts[cat] = cat_counts.get(cat, 0) + 1
|
||||
return selected
|
||||
|
||||
def _decay(turns: list[dict]) -> list[dict]:
|
||||
"""调用模块级 _apply_time_decay,传入当前 bridge_cfg。"""
|
||||
return _apply_time_decay(turns, cfg=bridge_cfg)
|
||||
|
||||
# ── 路由 ───────────────────────────────────────────
|
||||
|
||||
@app.get("/health")
|
||||
async def health() -> dict[str, str]:
|
||||
return {"status": "ok"}
|
||||
|
||||
@app.post("/memory/turn")
|
||||
async def add_turn(req: AddTurnRequest) -> dict[str, Any]:
|
||||
nonlocal _write_lock, _facts_write_lock
|
||||
if _write_lock is None:
|
||||
_write_lock = asyncio.Lock()
|
||||
if _facts_write_lock is None:
|
||||
_facts_write_lock = asyncio.Lock()
|
||||
|
||||
emb = _get_embedder()
|
||||
vec = emb.embed_query(req.content)
|
||||
from mem_bridge.compressor import count_tokens
|
||||
token_count = count_tokens(req.content)
|
||||
vec_array = np.array([vec], dtype=np.float32)
|
||||
|
||||
async with _write_lock:
|
||||
faiss_id = turn_store.next_faiss_id()
|
||||
id_array = np.array([faiss_id], dtype=np.int64)
|
||||
if not turns_index.is_trained:
|
||||
turns_index.build(vec_array, id_array)
|
||||
else:
|
||||
turns_index.add(vec_array, id_array)
|
||||
turns_index.save()
|
||||
turn_store.add_turn(
|
||||
req.session_id, req.role, req.content,
|
||||
faiss_id=faiss_id, token_count=token_count,
|
||||
)
|
||||
|
||||
# 异步触发摘要归档
|
||||
if context_mgr is not None:
|
||||
_safe_create_task(context_mgr.check_and_summarize(req.session_id))
|
||||
|
||||
# 异步触发事实提取(user+assistant 均触发)
|
||||
if fact_extractor is not None:
|
||||
all_turns = turn_store.list_turns(req.session_id)
|
||||
_safe_create_task(
|
||||
fact_extractor.extract_and_store(req.session_id, all_turns, _facts_write_lock) # type: ignore[arg-type]
|
||||
)
|
||||
|
||||
return {"ok": True, "faiss_id": faiss_id}
|
||||
|
||||
@app.get("/memory/context")
|
||||
async def get_context(
|
||||
session_id: str, query: str, token_budget: int = 2000
|
||||
) -> dict[str, Any]:
|
||||
emb = _get_embedder()
|
||||
q_vec = emb.embed_query(query)
|
||||
comp = Compressor(token_budget=token_budget)
|
||||
|
||||
# ── 1. facts(语义 + BM25 混合,category top-k)
|
||||
facts: list[dict[str, Any]] = _hybrid_search_facts(q_vec, query)
|
||||
|
||||
# ── 2. working_memory(最近 N 轮,无条件注入)
|
||||
working_memory = turn_store.get_recent_turns(
|
||||
session_id, n=bridge_cfg.working_memory_turns
|
||||
)
|
||||
|
||||
# ── 3. knowledge(文档知识库)
|
||||
long_term: list[dict[str, Any]] = []
|
||||
try:
|
||||
doc_pipeline._ensure_index_loaded()
|
||||
sem_results = doc_pipeline.search(query, top_k=bridge_cfg.long_term_top_k * 2)
|
||||
# BM25 混合(MetaStore FTS5)
|
||||
if bridge_cfg.hybrid_search_enabled:
|
||||
try:
|
||||
bm25_doc = doc_pipeline._store.bm25_search(query, limit=bridge_cfg.long_term_top_k * 2)
|
||||
except Exception as e:
|
||||
logger.debug("BM25 搜索失败,退化为纯语义: %s", e)
|
||||
bm25_doc = []
|
||||
sem_ids_doc = np.array([r.get("faiss_id", -1) for r in sem_results], dtype=np.int64)
|
||||
sem_dists_doc = np.array([r.get("score", 0.0) for r in sem_results], dtype=np.float32)
|
||||
scores_doc = _hybrid_scores(
|
||||
sem_ids_doc, sem_dists_doc, bm25_doc,
|
||||
bridge_cfg.hybrid_search_bm25_weight,
|
||||
)
|
||||
for r in sem_results:
|
||||
r["score"] = scores_doc.get(r.get("faiss_id", -1), r.get("score", 0.0))
|
||||
long_term = [
|
||||
{"chunk": r["chunk"], "path": r["path"], "score": r["score"]}
|
||||
for r in sorted(sem_results, key=lambda x: x.get("score", 0.0), reverse=True)[
|
||||
: bridge_cfg.long_term_top_k
|
||||
]
|
||||
]
|
||||
except RuntimeError:
|
||||
pass
|
||||
|
||||
# ── 4. short_term(summary + history)
|
||||
short_term: list[dict[str, Any]] = []
|
||||
all_session_turns = turn_store.list_turns(session_id)
|
||||
|
||||
# summary 轮次
|
||||
summary_turns = [t for t in all_session_turns if t["role"] == "summary"]
|
||||
if summary_turns:
|
||||
latest = summary_turns[-1]
|
||||
short_term.append({"role": "summary", "content": latest["content"], "score": 1.0})
|
||||
|
||||
# history:语义搜索 + BM25 混合 + 时间衰减
|
||||
session_faiss_ids = turn_store.get_faiss_ids(session_id)
|
||||
if session_faiss_ids:
|
||||
try:
|
||||
ids, dists = turns_index.search(q_vec, top_k=bridge_cfg.short_term_top_k * 2)
|
||||
valid_fids = set(session_faiss_ids)
|
||||
turns_by_fid = {
|
||||
t["faiss_id"]: t for t in all_session_turns
|
||||
if t.get("faiss_id") is not None
|
||||
}
|
||||
sem_hist: list[dict[str, Any]] = []
|
||||
for fid, dist in zip(ids, dists):
|
||||
if fid in valid_fids and fid in turns_by_fid:
|
||||
t = turns_by_fid[fid]
|
||||
if t.get("role") != "summary":
|
||||
sem_hist.append({
|
||||
"id": t["id"],
|
||||
"role": t["role"],
|
||||
"content": t["content"],
|
||||
"score": float(dist),
|
||||
"created_at": t.get("created_at", 0.0),
|
||||
})
|
||||
|
||||
# BM25 混合
|
||||
if bridge_cfg.hybrid_search_enabled:
|
||||
try:
|
||||
bm25_turns = turn_store.bm25_search(
|
||||
session_id, query, limit=bridge_cfg.short_term_top_k * 2
|
||||
)
|
||||
sem_ids_h = np.array([t.get("faiss_id", -1) for t in sem_hist], dtype=np.int64)
|
||||
sem_dists_h = np.array([t.get("score", 0.0) for t in sem_hist], dtype=np.float32)
|
||||
scores_h = _hybrid_scores(
|
||||
sem_ids_h, sem_dists_h, bm25_turns,
|
||||
bridge_cfg.hybrid_search_bm25_weight,
|
||||
)
|
||||
for t in sem_hist:
|
||||
fid = t.get("faiss_id") or turns_by_fid.get(t["id"], {}).get("faiss_id")
|
||||
if fid is not None:
|
||||
t["score"] = scores_h.get(int(fid), t.get("score", 0.0))
|
||||
except Exception as e:
|
||||
logger.debug("BM25 搜索失败,退化为纯语义: %s", e)
|
||||
|
||||
# 时间衰减
|
||||
sem_hist = _decay(sem_hist)
|
||||
short_term.extend(sem_hist[: bridge_cfg.short_term_top_k])
|
||||
|
||||
except Exception:
|
||||
recent = turn_store.get_recent_turns(session_id, n=bridge_cfg.recent_turns_verbatim)
|
||||
short_term.extend(
|
||||
{"role": t["role"], "content": t["content"], "score": 1.0,
|
||||
"id": t.get("id"), "created_at": t.get("created_at", 0.0)}
|
||||
for t in recent
|
||||
)
|
||||
|
||||
result = comp.compress(
|
||||
facts=facts,
|
||||
working_memory=working_memory,
|
||||
long_term=long_term,
|
||||
short_term=short_term,
|
||||
)
|
||||
return {
|
||||
"facts": result.facts,
|
||||
"long_term": result.long_term,
|
||||
"short_term": result.short_term,
|
||||
"formatted": result.formatted,
|
||||
"tokens_used": result.tokens_used,
|
||||
"budget": result.budget,
|
||||
}
|
||||
|
||||
@app.delete("/memory/session/{session_id}")
|
||||
async def delete_session(session_id: str) -> dict[str, bool]:
|
||||
turn_store.delete_session(session_id)
|
||||
return {"ok": True}
|
||||
|
||||
@app.get("/memory/search")
|
||||
async def search(q: str, limit: int = 5) -> dict[str, Any]:
|
||||
try:
|
||||
results = doc_pipeline.search(q, top_k=limit)
|
||||
return {"results": results}
|
||||
except RuntimeError:
|
||||
return {"results": []}
|
||||
|
||||
@app.get("/memory/facts")
|
||||
async def list_facts(
|
||||
category: str | None = None,
|
||||
entity: str | None = None,
|
||||
) -> dict[str, Any]:
|
||||
facts = fact_store.list_facts(category=category, entity=entity)
|
||||
return {"facts": facts}
|
||||
|
||||
@app.post("/memory/fact")
|
||||
async def add_fact_manual(req: AddFactRequest) -> dict[str, Any]:
|
||||
nonlocal _facts_write_lock
|
||||
if _facts_write_lock is None:
|
||||
_facts_write_lock = asyncio.Lock()
|
||||
emb = _get_embedder()
|
||||
vec = np.array(emb.embed_query(req.content), dtype=np.float32)
|
||||
async with _facts_write_lock:
|
||||
fact_id = fact_store.add_fact(
|
||||
req.content, req.category, req.entity, req.tags,
|
||||
"manual", vec,
|
||||
)
|
||||
return {"ok": True, "id": fact_id}
|
||||
|
||||
@app.delete("/memory/fact/{fact_id}")
|
||||
async def delete_fact(fact_id: int) -> dict[str, bool]:
|
||||
fact_store.delete_fact(fact_id)
|
||||
return {"ok": True}
|
||||
|
||||
# ── 管理端点 ──
|
||||
|
||||
@app.get("/admin/sessions")
|
||||
async def admin_list_sessions() -> dict[str, Any]:
|
||||
"""列出所有活跃会话(含轮次数和最近活跃时间)。"""
|
||||
sessions = turn_store.list_sessions()
|
||||
return {"sessions": sessions, "total": len(sessions)}
|
||||
|
||||
@app.delete("/admin/sessions/{session_id}")
|
||||
async def admin_delete_session(session_id: str) -> dict[str, bool]:
|
||||
"""清除指定会话的全部记忆(轮次 + 事实)。"""
|
||||
turn_store.delete_session(session_id)
|
||||
return {"ok": True}
|
||||
|
||||
@app.get("/admin/stats")
|
||||
async def admin_stats() -> dict[str, Any]:
|
||||
"""记忆使用统计:会话数、总轮次数、事实数。"""
|
||||
sessions = turn_store.list_sessions()
|
||||
total_turns = turn_store.count_turns()
|
||||
facts = fact_store.list_facts()
|
||||
return {
|
||||
"session_count": len(sessions),
|
||||
"turn_count": total_turns,
|
||||
"fact_count": len(facts),
|
||||
}
|
||||
|
||||
return app
|
||||
@@ -1,37 +0,0 @@
|
||||
"""OpenAI 兼容 Pydantic 模型 + 内部 ContextResponse。"""
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Any
|
||||
from pydantic import BaseModel, Field
|
||||
|
||||
|
||||
class ChatMessage(BaseModel):
|
||||
role: str
|
||||
content: str
|
||||
|
||||
|
||||
class ChatRequest(BaseModel):
|
||||
model: str = "mem-bridge"
|
||||
messages: list[ChatMessage]
|
||||
stream: bool = False
|
||||
temperature: float | None = None
|
||||
max_tokens: int | None = None
|
||||
session_id: str | None = Field(default=None, alias="x_session_id")
|
||||
|
||||
model_config = {"populate_by_name": True}
|
||||
|
||||
|
||||
class ContextResponse(BaseModel):
|
||||
long_term: list[dict[str, Any]]
|
||||
short_term: list[dict[str, Any]]
|
||||
formatted: str
|
||||
tokens_used: int
|
||||
budget: dict[str, int]
|
||||
|
||||
|
||||
class ChatResponse(BaseModel):
|
||||
id: str
|
||||
object: str = "chat.completion"
|
||||
model: str
|
||||
choices: list[dict[str, Any]]
|
||||
usage: dict[str, int] = Field(default_factory=dict)
|
||||
@@ -1,79 +0,0 @@
|
||||
"""提示词优化:Query 改写 + System Prompt 模板管理。"""
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
from pathlib import Path
|
||||
from typing import TYPE_CHECKING
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from mem_bridge.backends import BaseBackend
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
REWRITE_SYSTEM = (
|
||||
"你是检索助手,将用户问题改写为更适合语义检索的规范表达,"
|
||||
"不超过50字,只输出改写结果,不要任何解释。"
|
||||
)
|
||||
|
||||
|
||||
class QueryRewriter:
|
||||
"""将用户问题改写为规范语义表达,提升检索召回率。"""
|
||||
|
||||
def __init__(self, backend: BaseBackend, enabled: bool = False) -> None:
|
||||
self._backend = backend
|
||||
self._enabled = enabled
|
||||
|
||||
async def rewrite(self, query: str) -> str:
|
||||
"""改写查询;禁用或失败时静默回退原始 query。"""
|
||||
if not self._enabled or not query.strip():
|
||||
return query
|
||||
from mem_bridge.models import ChatMessage
|
||||
|
||||
try:
|
||||
messages = [
|
||||
ChatMessage(role="system", content=REWRITE_SYSTEM),
|
||||
ChatMessage(role="user", content=query),
|
||||
]
|
||||
result = await self._backend.chat(messages, stream=False)
|
||||
rewritten = str(result).strip()
|
||||
return rewritten if rewritten else query
|
||||
except Exception as e:
|
||||
logger.warning("Query 改写失败,使用原始 query: %s", e)
|
||||
return query
|
||||
|
||||
|
||||
class PromptBuilder:
|
||||
"""System Prompt 模板管理:按序拼接模板 + 原始 system + 记忆上下文。"""
|
||||
|
||||
def __init__(self, templates: dict[str, str]) -> None:
|
||||
self._templates = templates # 有序 dict(Python 3.7+ 保证插入顺序)
|
||||
|
||||
def build(self, original_system: str, context_text: str) -> str:
|
||||
"""拼接最终 system prompt。
|
||||
|
||||
拼接顺序:
|
||||
1. 各模板文本(按 active 列表顺序)
|
||||
2. original_system(若非空)
|
||||
3. context_text(记忆注入,若非空)
|
||||
"""
|
||||
parts: list[str] = []
|
||||
for content in self._templates.values():
|
||||
if content.strip():
|
||||
parts.append(content.strip())
|
||||
if original_system.strip():
|
||||
parts.append(original_system.strip())
|
||||
if context_text.strip():
|
||||
parts.append(context_text.strip())
|
||||
return "\n\n".join(parts)
|
||||
|
||||
@staticmethod
|
||||
def load_templates(dir: Path, names: list[str]) -> dict[str, str]:
|
||||
"""从目录加载指定模板文件(.txt),不存在的静默跳过。"""
|
||||
templates: dict[str, str] = {}
|
||||
for name in names:
|
||||
path = dir / f"{name}.txt"
|
||||
if path.exists():
|
||||
templates[name] = path.read_text(encoding="utf-8")
|
||||
else:
|
||||
logger.warning("模板文件不存在,跳过: %s", path)
|
||||
return templates
|
||||
@@ -1,261 +0,0 @@
|
||||
"""Router Service:OpenAI 兼容代理 + 复杂度路由 + 记忆注入。"""
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import hashlib
|
||||
import logging
|
||||
import uuid
|
||||
from typing import Any
|
||||
|
||||
import httpx
|
||||
from fastapi import FastAPI, Request
|
||||
from fastapi.responses import StreamingResponse
|
||||
|
||||
from mem_bridge.backends import BaseBackend, make_backend
|
||||
from mem_bridge.complexity import ComplexityScorer
|
||||
from mem_bridge.config import BridgeConfig
|
||||
from mem_bridge.models import ChatRequest, ChatMessage
|
||||
from mem_bridge.prompt_optimizer import PromptBuilder, QueryRewriter
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
def _safe_create_task(coro) -> None:
|
||||
"""在已运行的事件循环中创建 task,无事件循环时静默丢弃(fire-and-forget)。
|
||||
|
||||
使用 get_running_loop() 替代已废弃的 get_event_loop(),兼容 Python 3.10+/3.12。
|
||||
"""
|
||||
try:
|
||||
loop = asyncio.get_running_loop()
|
||||
loop.create_task(coro)
|
||||
except RuntimeError:
|
||||
# 无运行中的事件循环(TestClient 同步环境或模块级调用)
|
||||
try:
|
||||
coro.close()
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
|
||||
def create_router_app(
|
||||
bridge_cfg: BridgeConfig,
|
||||
embedder=None,
|
||||
backends_override: dict[str, BaseBackend] | None = None,
|
||||
memory_base_url: str | None = "http://localhost:8001",
|
||||
prompt_builder: PromptBuilder | None = None,
|
||||
query_rewriter: QueryRewriter | None = None,
|
||||
) -> FastAPI:
|
||||
app = FastAPI(title="mem-bridge router-service")
|
||||
|
||||
# 初始化后端
|
||||
backends: dict[str, BaseBackend] = backends_override or {}
|
||||
if not backends_override:
|
||||
for name, cfg in bridge_cfg.backends.items():
|
||||
backends[name] = make_backend(name, cfg)
|
||||
|
||||
# 初始化复杂度评分器
|
||||
scorer: ComplexityScorer | None = None
|
||||
if embedder and bridge_cfg.complexity_templates_path.exists():
|
||||
templates = ComplexityScorer.load_templates(bridge_cfg.complexity_templates_path)
|
||||
scorer = ComplexityScorer(embedder, templates, threshold=bridge_cfg.complexity_threshold)
|
||||
logger.info("ComplexityScorer loaded, %d templates", len(templates))
|
||||
|
||||
# 初始化 PromptBuilder
|
||||
if prompt_builder is None:
|
||||
pt_templates = PromptBuilder.load_templates(
|
||||
bridge_cfg.prompt_templates_dir,
|
||||
bridge_cfg.prompt_templates_active,
|
||||
)
|
||||
prompt_builder = PromptBuilder(pt_templates)
|
||||
|
||||
# 初始化 QueryRewriter
|
||||
if query_rewriter is None:
|
||||
qr_backend_name = bridge_cfg.query_rewrite_backend or bridge_cfg.default_backend
|
||||
qr_bk: BaseBackend | None = backends.get(qr_backend_name) or (
|
||||
next(iter(backends.values())) if backends else None
|
||||
)
|
||||
if qr_bk is not None:
|
||||
query_rewriter = QueryRewriter(
|
||||
backend=qr_bk,
|
||||
enabled=bridge_cfg.query_rewrite_enabled,
|
||||
)
|
||||
else:
|
||||
# 无后端时创建禁用状态的 QueryRewriter(enabled=False 不会实际调用)
|
||||
|
||||
class _NullBackend(BaseBackend):
|
||||
pass
|
||||
|
||||
query_rewriter = QueryRewriter(backend=_NullBackend(), enabled=False)
|
||||
|
||||
def _select_backend(query: str, task_type: str = "") -> tuple[BaseBackend, list[str]]:
|
||||
"""Select primary backend and build fallback chain.
|
||||
|
||||
Returns (primary_backend, fallback_backend_names).
|
||||
"""
|
||||
# Task-type routing takes priority
|
||||
if task_type and task_type in bridge_cfg.routing_rules:
|
||||
name = bridge_cfg.routing_rules[task_type]
|
||||
if name in backends:
|
||||
chain = [n for n in bridge_cfg.fallback_chain if n != name and n in backends]
|
||||
return backends[name], chain
|
||||
|
||||
# Complexity-based routing
|
||||
if scorer and scorer.is_complex(query):
|
||||
name = bridge_cfg.heavy_backend or bridge_cfg.default_backend
|
||||
else:
|
||||
name = bridge_cfg.default_backend
|
||||
if not name or name not in backends:
|
||||
name = next(iter(backends))
|
||||
|
||||
chain = [n for n in bridge_cfg.fallback_chain if n != name and n in backends]
|
||||
return backends[name], chain
|
||||
|
||||
async def _call_with_fallback(
|
||||
primary: BaseBackend,
|
||||
fallback_names: list[str],
|
||||
messages: list[ChatMessage],
|
||||
stream: bool,
|
||||
temperature: float | None,
|
||||
max_tokens: int | None,
|
||||
):
|
||||
"""Call primary backend, fallback to chain on timeout/error."""
|
||||
timeout = bridge_cfg.fallback_timeout
|
||||
|
||||
# Try primary
|
||||
try:
|
||||
return await asyncio.wait_for(
|
||||
primary.chat(messages, stream=stream,
|
||||
temperature=temperature, max_tokens=max_tokens),
|
||||
timeout=timeout,
|
||||
)
|
||||
except (asyncio.TimeoutError, Exception) as e:
|
||||
logger.warning("Primary backend failed: %s — trying fallback chain", e)
|
||||
|
||||
# Try fallback chain
|
||||
for fb_name in fallback_names:
|
||||
fb = backends.get(fb_name)
|
||||
if fb is None:
|
||||
continue
|
||||
try:
|
||||
result = await asyncio.wait_for(
|
||||
fb.chat(messages, stream=stream,
|
||||
temperature=temperature, max_tokens=max_tokens),
|
||||
timeout=timeout,
|
||||
)
|
||||
logger.info("Fallback to %s succeeded", fb_name)
|
||||
return result
|
||||
except (asyncio.TimeoutError, Exception) as e:
|
||||
logger.warning("Fallback %s failed: %s", fb_name, e)
|
||||
|
||||
raise RuntimeError("All backends failed (primary + fallback chain)")
|
||||
|
||||
async def _fetch_context(session_id: str, query: str) -> str:
|
||||
if not memory_base_url:
|
||||
return ""
|
||||
try:
|
||||
async with httpx.AsyncClient(timeout=5.0) as client:
|
||||
resp = await client.get(
|
||||
f"{memory_base_url}/memory/context",
|
||||
params={"session_id": session_id, "query": query,
|
||||
"token_budget": bridge_cfg.token_budget},
|
||||
)
|
||||
if resp.status_code == 200:
|
||||
return resp.json().get("formatted", "")
|
||||
except Exception as e:
|
||||
logger.warning("memory-service 不可达: %s", e)
|
||||
return ""
|
||||
|
||||
async def _store_turn(session_id: str, role: str, content: str) -> None:
|
||||
if not memory_base_url:
|
||||
return
|
||||
try:
|
||||
async with httpx.AsyncClient(timeout=3.0) as client:
|
||||
await client.post(
|
||||
f"{memory_base_url}/memory/turn",
|
||||
json={"session_id": session_id, "role": role, "content": content},
|
||||
)
|
||||
except Exception:
|
||||
pass # fire-and-forget,失败静默
|
||||
|
||||
@app.get("/health")
|
||||
async def health():
|
||||
return {"status": "ok"}
|
||||
|
||||
@app.get("/v1/models")
|
||||
async def list_models():
|
||||
return {
|
||||
"object": "list",
|
||||
"data": [
|
||||
{"id": name, "object": "model", "owned_by": "mem-bridge"}
|
||||
for name in backends
|
||||
],
|
||||
}
|
||||
|
||||
@app.post("/v1/chat/completions")
|
||||
async def chat_completions(request: Request):
|
||||
body = await request.json()
|
||||
req = ChatRequest.model_validate(body)
|
||||
|
||||
# 提取 session_id
|
||||
client_host = request.client.host if request.client else "local"
|
||||
session_id = (
|
||||
request.headers.get("X-Session-Id")
|
||||
or req.session_id
|
||||
or hashlib.md5(f"{client_host}:{req.model}".encode()).hexdigest()[:16]
|
||||
)
|
||||
|
||||
# 提取 task_type(用于路由规则)
|
||||
task_type = request.headers.get("X-Task-Type", "")
|
||||
|
||||
query = req.messages[-1].content if req.messages else ""
|
||||
|
||||
# Query 改写(可选,enabled=False 时直接返回原始 query)
|
||||
query_for_retrieval = await query_rewriter.rewrite(query)
|
||||
|
||||
# 拉取记忆上下文(使用改写后的 query)
|
||||
context_text = await _fetch_context(session_id, query_for_retrieval)
|
||||
|
||||
# 构建最终 system prompt(模板 + 原始 system + 记忆上下文)
|
||||
messages = list(req.messages)
|
||||
original_system = next(
|
||||
(m.content for m in messages if m.role == "system"), ""
|
||||
)
|
||||
final_system = prompt_builder.build(original_system, context_text)
|
||||
messages = [m for m in messages if m.role != "system"]
|
||||
if final_system:
|
||||
messages.insert(0, ChatMessage(role="system", content=final_system))
|
||||
|
||||
# 选择后端 + fallback 链
|
||||
backend, fallback_names = _select_backend(query, task_type)
|
||||
|
||||
# 存储用户消息(fire-and-forget)
|
||||
_safe_create_task(_store_turn(session_id, "user", query))
|
||||
|
||||
if req.stream:
|
||||
async def stream_gen():
|
||||
full_response: list[str] = []
|
||||
gen = await _call_with_fallback(
|
||||
backend, fallback_names, messages, stream=True,
|
||||
temperature=req.temperature, max_tokens=req.max_tokens,
|
||||
)
|
||||
async for chunk in gen:
|
||||
full_response.append(chunk)
|
||||
yield f"data: {chunk}\n\n"
|
||||
yield "data: [DONE]\n\n"
|
||||
_safe_create_task(_store_turn(session_id, "assistant", "".join(full_response)))
|
||||
return StreamingResponse(stream_gen(), media_type="text/event-stream")
|
||||
|
||||
content = await _call_with_fallback(
|
||||
backend, fallback_names, messages, stream=False,
|
||||
temperature=req.temperature, max_tokens=req.max_tokens,
|
||||
)
|
||||
_safe_create_task(_store_turn(session_id, "assistant", str(content)))
|
||||
|
||||
return {
|
||||
"id": f"chatcmpl-{uuid.uuid4().hex[:8]}",
|
||||
"object": "chat.completion",
|
||||
"model": req.model,
|
||||
"choices": [{"index": 0, "message": {"role": "assistant", "content": content}, "finish_reason": "stop"}],
|
||||
"usage": {},
|
||||
}
|
||||
|
||||
return app
|
||||
@@ -1,168 +0,0 @@
|
||||
"""对话轮次 SQLite 存储。"""
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
import sqlite3
|
||||
import time
|
||||
from pathlib import Path
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
_DDL = """
|
||||
CREATE TABLE IF NOT EXISTS turns (
|
||||
id INTEGER PRIMARY KEY AUTOINCREMENT,
|
||||
session_id TEXT NOT NULL,
|
||||
role TEXT NOT NULL,
|
||||
content TEXT NOT NULL,
|
||||
faiss_id INTEGER,
|
||||
token_count INTEGER DEFAULT 0,
|
||||
created_at REAL NOT NULL
|
||||
);
|
||||
CREATE TABLE IF NOT EXISTS sessions (
|
||||
session_id TEXT PRIMARY KEY,
|
||||
created_at REAL NOT NULL,
|
||||
last_active REAL NOT NULL
|
||||
);
|
||||
CREATE INDEX IF NOT EXISTS idx_turns_session ON turns(session_id, created_at);
|
||||
"""
|
||||
|
||||
_FTS_DDL = """
|
||||
CREATE VIRTUAL TABLE IF NOT EXISTS turns_fts USING fts5(
|
||||
content, role, session_id,
|
||||
content='turns', content_rowid='id'
|
||||
);
|
||||
CREATE TRIGGER IF NOT EXISTS turns_ai AFTER INSERT ON turns BEGIN
|
||||
INSERT INTO turns_fts(rowid, content, role, session_id)
|
||||
VALUES (new.id, new.content, new.role, new.session_id);
|
||||
END;
|
||||
CREATE TRIGGER IF NOT EXISTS turns_ad AFTER DELETE ON turns BEGIN
|
||||
INSERT INTO turns_fts(turns_fts, rowid, content, role, session_id)
|
||||
VALUES ('delete', old.id, old.content, old.role, old.session_id);
|
||||
END;
|
||||
"""
|
||||
|
||||
|
||||
class TurnStore:
|
||||
def __init__(self, db_path: Path) -> None:
|
||||
self._db_path = db_path
|
||||
self._init()
|
||||
|
||||
def _connect(self) -> sqlite3.Connection:
|
||||
conn = sqlite3.connect(self._db_path)
|
||||
conn.row_factory = sqlite3.Row
|
||||
conn.execute("PRAGMA journal_mode=WAL;")
|
||||
conn.execute("PRAGMA synchronous=NORMAL;")
|
||||
return conn
|
||||
|
||||
def _init(self) -> None:
|
||||
with self._connect() as conn:
|
||||
conn.executescript(_DDL)
|
||||
conn.executescript(_FTS_DDL)
|
||||
|
||||
def add_turn(
|
||||
self, session_id: str, role: str, content: str,
|
||||
faiss_id: int | None = None, token_count: int = 0,
|
||||
) -> int:
|
||||
now = time.time()
|
||||
with self._connect() as conn:
|
||||
conn.execute(
|
||||
"INSERT OR IGNORE INTO sessions VALUES (?,?,?)",
|
||||
(session_id, now, now),
|
||||
)
|
||||
conn.execute(
|
||||
"UPDATE sessions SET last_active=? WHERE session_id=?",
|
||||
(now, session_id),
|
||||
)
|
||||
cur = conn.execute(
|
||||
"INSERT INTO turns(session_id,role,content,faiss_id,token_count,created_at)"
|
||||
" VALUES (?,?,?,?,?,?)",
|
||||
(session_id, role, content, faiss_id, token_count, now),
|
||||
)
|
||||
return cur.lastrowid # type: ignore[return-value]
|
||||
|
||||
def list_turns(self, session_id: str) -> list[dict]:
|
||||
with self._connect() as conn:
|
||||
rows = conn.execute(
|
||||
"SELECT * FROM turns WHERE session_id=? ORDER BY created_at",
|
||||
(session_id,),
|
||||
).fetchall()
|
||||
return [dict(r) for r in rows]
|
||||
|
||||
def get_recent_turns(self, session_id: str, n: int = 2) -> list[dict]:
|
||||
with self._connect() as conn:
|
||||
rows = conn.execute(
|
||||
"SELECT * FROM turns WHERE session_id=? ORDER BY created_at DESC LIMIT ?",
|
||||
(session_id, n),
|
||||
).fetchall()
|
||||
return [dict(r) for r in reversed(rows)]
|
||||
|
||||
def get_faiss_ids(self, session_id: str) -> list[int]:
|
||||
with self._connect() as conn:
|
||||
rows = conn.execute(
|
||||
"SELECT faiss_id FROM turns WHERE session_id=? AND faiss_id IS NOT NULL",
|
||||
(session_id,),
|
||||
).fetchall()
|
||||
return [r[0] for r in rows]
|
||||
|
||||
def next_faiss_id(self) -> int:
|
||||
with self._connect() as conn:
|
||||
row = conn.execute("SELECT MAX(faiss_id) FROM turns").fetchone()
|
||||
return (row[0] + 1) if row[0] is not None else 0
|
||||
|
||||
def list_sessions(self) -> list[dict]:
|
||||
"""返回所有活跃会话列表,含最近活跃时间和轮次数。"""
|
||||
with self._connect() as conn:
|
||||
rows = conn.execute(
|
||||
"SELECT s.session_id, s.created_at, s.last_active, "
|
||||
" COUNT(t.id) AS turn_count "
|
||||
"FROM sessions s "
|
||||
"LEFT JOIN turns t ON t.session_id = s.session_id "
|
||||
"GROUP BY s.session_id "
|
||||
"ORDER BY s.last_active DESC",
|
||||
).fetchall()
|
||||
return [dict(r) for r in rows]
|
||||
|
||||
def count_turns(self) -> int:
|
||||
"""返回所有会话的总轮次数(用于统计)。"""
|
||||
with self._connect() as conn:
|
||||
row = conn.execute("SELECT COUNT(*) FROM turns").fetchone()
|
||||
return row[0] if row else 0
|
||||
|
||||
def delete_session(self, session_id: str) -> None:
|
||||
with self._connect() as conn:
|
||||
conn.execute("DELETE FROM turns WHERE session_id=?", (session_id,))
|
||||
conn.execute("DELETE FROM sessions WHERE session_id=?", (session_id,))
|
||||
|
||||
def delete_turns_by_ids(self, ids: list[int]) -> None:
|
||||
"""批量删除指定 id 的轮次(摘要归档时清理旧数据)。"""
|
||||
if not ids:
|
||||
return
|
||||
placeholders = ",".join("?" * len(ids))
|
||||
with self._connect() as conn:
|
||||
conn.execute(f"DELETE FROM turns WHERE id IN ({placeholders})", ids)
|
||||
|
||||
def bm25_search(self, session_id: str, query: str, limit: int = 20) -> list[dict]:
|
||||
"""FTS5 BM25 全文检索指定 session 内的对话轮次。
|
||||
|
||||
bm25_score 为负值,越负越相关。
|
||||
空查询直接返回空列表。
|
||||
"""
|
||||
tokens = [t for t in query.split() if t]
|
||||
if not tokens:
|
||||
return []
|
||||
fts_query = " OR ".join(tokens)
|
||||
with self._connect() as conn:
|
||||
rows = conn.execute(
|
||||
"""
|
||||
SELECT t.id, t.session_id, t.role, t.content,
|
||||
t.faiss_id, t.token_count, t.created_at,
|
||||
bm25(turns_fts) AS bm25_score
|
||||
FROM turns_fts
|
||||
JOIN turns t ON t.id = turns_fts.rowid
|
||||
WHERE turns_fts MATCH ? AND t.session_id = ?
|
||||
ORDER BY bm25_score
|
||||
LIMIT ?
|
||||
""",
|
||||
(fts_query, session_id, limit),
|
||||
).fetchall()
|
||||
return [dict(r) for r in rows]
|
||||
@@ -1,27 +0,0 @@
|
||||
# RK3588 设备端依赖
|
||||
# 系统已预装:Python 3.12.3、numpy 2.4.2、opencv-python 4.13.0
|
||||
|
||||
# 推理(需手动安装 wheel)
|
||||
# rknn-toolkit-lite2==2.3.2
|
||||
|
||||
# 推理后端(RKNN 优先;fastembed 作为回退)
|
||||
fastembed>=0.4.0 # 回退后端:ONNX 推理,无需 PyTorch
|
||||
tokenizers>=0.15.0 # ONNX 后端分词器
|
||||
|
||||
# 向量索引
|
||||
faiss-cpu>=1.7.4
|
||||
|
||||
# 文档解析
|
||||
pymupdf>=1.23.0
|
||||
python-docx>=1.1.0
|
||||
openpyxl>=3.1.0
|
||||
pytesseract>=0.3.10
|
||||
|
||||
# mem-bridge 中间层服务
|
||||
fastapi>=0.111.0
|
||||
uvicorn[standard]>=0.30.0
|
||||
httpx>=0.27.0
|
||||
pyyaml>=6.0.1
|
||||
tiktoken>=0.7.0
|
||||
anthropic>=0.28.0
|
||||
openai>=1.35.0
|
||||
-95
@@ -1,95 +0,0 @@
|
||||
"""mem-bridge 启动入口:单进程启动 memory-service 和 router-service。"""
|
||||
from __future__ import annotations
|
||||
|
||||
import argparse
|
||||
import logging
|
||||
import os
|
||||
import sys
|
||||
import threading
|
||||
from pathlib import Path
|
||||
|
||||
import uvicorn
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
def main() -> None:
|
||||
parser = argparse.ArgumentParser(description="mem-bridge 服务启动")
|
||||
parser.add_argument("--config", default="config.yaml", help="配置文件路径")
|
||||
parser.add_argument(
|
||||
"--service", choices=["memory", "router", "both"], default="both",
|
||||
help="启动哪个服务(默认 both)",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--debug", action="store_true",
|
||||
default=os.environ.get("KVM_DEBUG", "").lower() in ("1", "true"),
|
||||
help="Enable debug logging (env: KVM_DEBUG=1)",
|
||||
)
|
||||
args = parser.parse_args()
|
||||
|
||||
is_tty = sys.stderr.isatty()
|
||||
logging.basicConfig(
|
||||
format="%(asctime)s %(levelname)s [%(name)s]: %(message)s" if is_tty
|
||||
else "[%(name)s] %(levelname)s: %(message)s",
|
||||
datefmt="%H:%M:%S" if is_tty else None,
|
||||
level=logging.DEBUG if args.debug else logging.INFO,
|
||||
)
|
||||
|
||||
cfg_path = Path(args.config)
|
||||
if not cfg_path.exists():
|
||||
logger.error("配置文件不存在: %s", cfg_path)
|
||||
raise SystemExit(1)
|
||||
|
||||
from mem_bridge.config import BridgeConfig
|
||||
bridge_cfg = BridgeConfig.from_yaml(cfg_path)
|
||||
|
||||
from embed_db.config import Config as EmbedConfig
|
||||
embed_cfg = EmbedConfig(db_dir=bridge_cfg.db_dir)
|
||||
|
||||
if args.service == "memory":
|
||||
from mem_bridge.memory_service import create_memory_app
|
||||
app = create_memory_app(bridge_cfg, embed_cfg)
|
||||
uvicorn.run(app, host=bridge_cfg.memory_host, port=bridge_cfg.memory_port)
|
||||
return
|
||||
|
||||
if args.service == "router":
|
||||
from mem_bridge.router_service import create_router_app
|
||||
app = create_router_app(
|
||||
bridge_cfg,
|
||||
memory_base_url=f"http://localhost:{bridge_cfg.memory_port}",
|
||||
)
|
||||
uvicorn.run(app, host=bridge_cfg.router_host, port=bridge_cfg.router_port)
|
||||
return
|
||||
|
||||
# both: 守护线程运行 memory-service,主线程运行 router-service
|
||||
from mem_bridge.memory_service import create_memory_app
|
||||
from mem_bridge.router_service import create_router_app
|
||||
|
||||
memory_app = create_memory_app(bridge_cfg, embed_cfg)
|
||||
router_app = create_router_app(
|
||||
bridge_cfg,
|
||||
memory_base_url=f"http://localhost:{bridge_cfg.memory_port}",
|
||||
)
|
||||
|
||||
def run_memory() -> None:
|
||||
uvicorn.run(
|
||||
memory_app,
|
||||
host=bridge_cfg.memory_host,
|
||||
port=bridge_cfg.memory_port,
|
||||
log_level="debug" if args.debug else "info",
|
||||
)
|
||||
|
||||
t = threading.Thread(target=run_memory, daemon=True)
|
||||
t.start()
|
||||
logger.info("memory-service 启动中: http://localhost:%d", bridge_cfg.memory_port)
|
||||
|
||||
uvicorn.run(
|
||||
router_app,
|
||||
host=bridge_cfg.router_host,
|
||||
port=bridge_cfg.router_port,
|
||||
log_level="debug" if args.debug else "info",
|
||||
)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
Vendored
+6
-5
@@ -1,11 +1,12 @@
|
||||
Package: kvm-meta
|
||||
Version: 1.0.0-1
|
||||
Version: 2.0.0-1
|
||||
Architecture: all
|
||||
Maintainer: KVM-Privacy <noreply@kvm-privacy.local>
|
||||
Depends: kvm-server (>= 1.0.0), kvm-agent (>= 1.0.0), kvm-privacy (>= 1.0.0)
|
||||
Recommends: kvm-bridge, kvm-npu, kvm-rkllm
|
||||
Depends: kvm-server (>= 1.0.0), kvm-agent (>= 2.0.0), kvm-privacy (>= 1.0.0)
|
||||
Recommends: kvm-bridge, kvm-npu, kvm-rkllm, kvm-mitm
|
||||
Description: KVM-Privacy complete product (metapackage)
|
||||
Installs the full KVM-Privacy product stack:
|
||||
kvm-server (Go KVM core), kvm-agent (AI automation),
|
||||
kvm-privacy (PII detection), kvm-npu (NPU inference),
|
||||
kvm-rkllm (local LLM), and optionally kvm-bridge (memory).
|
||||
kvm-privacy (PII detection), kvm-mitm (privacy gateway),
|
||||
kvm-npu (NPU inference), kvm-rkllm (local LLM),
|
||||
and optionally kvm-bridge (session memory).
|
||||
|
||||
Vendored
+7
-7
@@ -1,10 +1,10 @@
|
||||
Package: kvm-mitm
|
||||
Version: 1.0.0-2
|
||||
Architecture: all
|
||||
Version: 2.0.0-1
|
||||
Architecture: arm64
|
||||
Maintainer: KVM-Privacy <noreply@kvm-privacy.local>
|
||||
Depends: python3 (>= 3.9)
|
||||
Recommends: dnsmasq, iptables, python3-aiohttp, python3-httpx
|
||||
Description: KVM Privacy Gateway + Document Processor
|
||||
Privacy MITM interceptor (regular + transparent mode),
|
||||
document privacy processor, and LAN gateway scripts.
|
||||
Depends: libc6
|
||||
Recommends: dnsmasq, iptables
|
||||
Description: KVM Privacy Gateway (Rust)
|
||||
Privacy MITM interceptor using hudsucker transparent proxy.
|
||||
Scans uploads for PII using info-privacy-rs.
|
||||
Includes LAN gateway scripts for NAT + transparent proxy.
|
||||
|
||||
Vendored
+11
-15
@@ -2,29 +2,25 @@
|
||||
set -e
|
||||
case "$1" in
|
||||
configure)
|
||||
mkdir -p /etc/kvm-privacy
|
||||
if [ ! -f /etc/kvm-privacy/secrets.env ]; then
|
||||
cat > /etc/kvm-privacy/secrets.env <<'EOF'
|
||||
KVM_MITM_DB_HOST=localhost
|
||||
KVM_MITM_DB_USER=kvm_mitm
|
||||
KVM_MITM_DB_PASS=changeme
|
||||
KVM_MITM_DB_NAME=kvm
|
||||
EOF
|
||||
chmod 600 /etc/kvm-privacy/secrets.env
|
||||
chown root:root /etc/kvm-privacy/secrets.env
|
||||
echo "WARNING: Edit /etc/kvm-privacy/secrets.env with actual DB password"
|
||||
fi
|
||||
mkdir -p /etc/kvm-mitm
|
||||
|
||||
# CA certificate directory
|
||||
mkdir -p /etc/kvm-privacy/ca
|
||||
chmod 700 /etc/kvm-privacy/ca
|
||||
|
||||
# State directory for privacy mode persistence
|
||||
mkdir -p /var/lib/kvm-privacy
|
||||
chown root:root /var/lib/kvm-privacy
|
||||
|
||||
# Install dnsmasq config if not present
|
||||
mkdir -p /etc/kvm
|
||||
if [ ! -f /etc/kvm/dnsmasq-lan.conf ]; then
|
||||
cp /usr/lib/kvm-mitm/network/dnsmasq-lan.conf /etc/kvm/dnsmasq-lan.conf
|
||||
cp /usr/lib/kvm-mitm/network/dnsmasq-lan.conf /etc/kvm/dnsmasq-lan.conf 2>/dev/null || true
|
||||
fi
|
||||
|
||||
systemctl daemon-reload
|
||||
systemctl enable kvm-mitm.service 2>/dev/null || true
|
||||
systemctl enable doc-processor.service 2>/dev/null || true
|
||||
# gateway disabled by default — user must enable in kvm.toml
|
||||
systemctl start kvm-mitm.service 2>/dev/null || true
|
||||
;;
|
||||
esac
|
||||
exit 0
|
||||
|
||||
Vendored
+3
-14
@@ -2,21 +2,10 @@
|
||||
set -e
|
||||
case "$1" in
|
||||
purge)
|
||||
# Clean runtime-generated __pycache__ files
|
||||
rm -rf /usr/lib/kvm-mitm/privacy_gateway/__pycache__
|
||||
rm -rf /usr/lib/kvm-mitm/privacy_gateway/tests/__pycache__
|
||||
rm -rf /usr/lib/kvm-mitm/doc_processor/__pycache__
|
||||
rm -rf /usr/lib/kvm-mitm/kvm_common/__pycache__
|
||||
rm -rf /usr/lib/kvm-mitm/network
|
||||
rmdir --ignore-fail-on-non-empty /usr/lib/kvm-mitm/privacy_gateway/tests 2>/dev/null || true
|
||||
rmdir --ignore-fail-on-non-empty /usr/lib/kvm-mitm/privacy_gateway 2>/dev/null || true
|
||||
rmdir --ignore-fail-on-non-empty /usr/lib/kvm-mitm/doc_processor/static 2>/dev/null || true
|
||||
rmdir --ignore-fail-on-non-empty /usr/lib/kvm-mitm/doc_processor 2>/dev/null || true
|
||||
rmdir --ignore-fail-on-non-empty /usr/lib/kvm-mitm/kvm_common 2>/dev/null || true
|
||||
rmdir --ignore-fail-on-non-empty /usr/lib/kvm-mitm 2>/dev/null || true
|
||||
# Clean config directory
|
||||
rm -rf /etc/kvm-privacy
|
||||
rm -rf /usr/lib/kvm-mitm
|
||||
rm -rf /etc/kvm-mitm
|
||||
rm -f /etc/kvm/dnsmasq-lan.conf
|
||||
rm -rf /var/lib/kvm-privacy
|
||||
;;
|
||||
esac
|
||||
exit 0
|
||||
|
||||
Vendored
-2
@@ -4,8 +4,6 @@ case "$1" in
|
||||
remove|purge)
|
||||
systemctl stop kvm-mitm.service 2>/dev/null || true
|
||||
systemctl disable kvm-mitm.service 2>/dev/null || true
|
||||
systemctl stop doc-processor.service 2>/dev/null || true
|
||||
systemctl disable doc-processor.service 2>/dev/null || true
|
||||
systemctl stop kvm-gateway.service 2>/dev/null || true
|
||||
systemctl disable kvm-gateway.service 2>/dev/null || true
|
||||
;;
|
||||
|
||||
+1
-1
@@ -2,7 +2,7 @@
|
||||
Description=KVM LAN Gateway (NAT + Transparent Proxy)
|
||||
After=network-online.target
|
||||
Wants=network-online.target
|
||||
Before=privacy-gateway.service
|
||||
Before=kvm-mitm.service
|
||||
|
||||
[Service]
|
||||
Type=oneshot
|
||||
|
||||
+13
-5
@@ -1,6 +1,6 @@
|
||||
[Unit]
|
||||
Description=KVM Privacy MITM Interceptor
|
||||
After=network.target mariadb.service info-privacy.service
|
||||
Description=KVM Privacy Gateway (Rust, hudsucker proxy + axum API)
|
||||
After=network.target info-privacy.service kvm-gateway.service
|
||||
Wants=info-privacy.service
|
||||
StartLimitInterval=300
|
||||
StartLimitBurst=5
|
||||
@@ -8,11 +8,19 @@ StartLimitBurst=5
|
||||
[Service]
|
||||
Type=simple
|
||||
User=root
|
||||
EnvironmentFile=/etc/kvm-privacy/secrets.env
|
||||
ExecStart=/usr/bin/kvm-mitm
|
||||
EnvironmentFile=-/etc/kvm-mitm/secrets.env
|
||||
Environment=PROXY_PORT=8888
|
||||
Environment=API_PORT=8889
|
||||
Environment=INFO_PRIVACY_URL=http://localhost:8001
|
||||
Environment=RKLLM_URL=http://localhost:8891
|
||||
Environment=STATE_FILE=/var/lib/kvm-privacy/state.json
|
||||
Environment=CA_DIR=/etc/kvm-privacy/ca
|
||||
RuntimeDirectory=kvm-mitm
|
||||
RuntimeDirectoryMode=0750
|
||||
ExecStart=/usr/sbin/privacy-gateway-rs
|
||||
Restart=on-failure
|
||||
RestartSec=10
|
||||
MemoryMax=768M
|
||||
MemoryMax=256M
|
||||
LimitNOFILE=8192
|
||||
StandardOutput=journal
|
||||
StandardError=journal
|
||||
|
||||
@@ -1,26 +0,0 @@
|
||||
[Unit]
|
||||
Description=KVM Privacy Gateway (mitmproxy transparent)
|
||||
After=network.target info-privacy.service
|
||||
Wants=network.target info-privacy.service
|
||||
StartLimitInterval=300
|
||||
StartLimitBurst=5
|
||||
|
||||
[Service]
|
||||
Type=simple
|
||||
User=root
|
||||
ExecStart=/usr/local/bin/mitmdump \
|
||||
--mode transparent \
|
||||
--listen-host 0.0.0.0 \
|
||||
--listen-port 8888 \
|
||||
--ssl-insecure \
|
||||
-s /usr/lib/kvm-mitm/privacy_gateway/addon.py
|
||||
Environment=PYTHONPATH=/usr/lib/kvm-mitm
|
||||
Restart=on-failure
|
||||
RestartSec=10
|
||||
MemoryMax=768M
|
||||
LimitNOFILE=8192
|
||||
StandardOutput=journal
|
||||
StandardError=journal
|
||||
|
||||
[Install]
|
||||
WantedBy=multi-user.target
|
||||
Vendored
-3
@@ -1,3 +0,0 @@
|
||||
#!/bin/bash
|
||||
export PYTHONPATH=/usr/lib/kvm-mitm
|
||||
exec python3 -m privacy_gateway.mitm_launcher "$@"
|
||||
Vendored
+1
@@ -0,0 +1 @@
|
||||
/etc/npu-daemon/config.yaml
|
||||
Vendored
+1
@@ -2,6 +2,7 @@ Package: kvm-npu
|
||||
Version: 1.0.0-1
|
||||
Architecture: arm64
|
||||
Maintainer: KVM-Privacy <noreply@kvm-privacy.local>
|
||||
Depends: libc6 (>= 2.31)
|
||||
Description: NPU Daemon - Centralized RKNN inference (OCR, face, redaction)
|
||||
Manages RKNN NPU cores on RK3588 for OCR detection,
|
||||
recognition, face detection, and privacy redaction.
|
||||
|
||||
+8
@@ -0,0 +1,8 @@
|
||||
#!/bin/bash
|
||||
set -e
|
||||
case "$1" in
|
||||
purge)
|
||||
rm -rf /etc/npu-daemon
|
||||
;;
|
||||
esac
|
||||
exit 0
|
||||
+18
-10
@@ -1,20 +1,28 @@
|
||||
server:
|
||||
host: 0.0.0.0
|
||||
host: "0.0.0.0"
|
||||
port: 8004
|
||||
|
||||
models:
|
||||
ocr_det: /usr/share/kvm-npu/models/ocr/ppocrv4_det.rknn
|
||||
ocr_rec_ch: /usr/share/kvm-npu/models/ocr/ch/ppocrv4_rec.rknn
|
||||
ocr_rec_en: /usr/share/kvm-npu/models/ocr/en/ppocrv4_rec.rknn
|
||||
ocr_dict_ch: /usr/share/kvm-npu/models/ocr/ch/ppocr_keys.txt
|
||||
ocr_dict_en: /usr/share/kvm-npu/models/ocr/en/ppocr_keys.txt
|
||||
face_det: /usr/share/kvm-npu/models/face/face_detection_short_range_rk3588.rknn
|
||||
# PP-OCRv4 detection model (text region detection)
|
||||
ocr_det: /data/project/KVM-privacy/deps/KVM/models/ocr/ppocrv4_det.rknn
|
||||
# PP-OCRv4 recognition model (Chinese)
|
||||
ocr_rec_ch: /data/project/KVM-privacy/deps/KVM/models/ocr/ch/ppocrv4_rec.rknn
|
||||
# PP-OCRv4 recognition model (English, optional)
|
||||
ocr_rec_en: /data/project/KVM-privacy/deps/KVM/models/ocr/en/ppocrv4_rec.rknn
|
||||
# Character dictionaries
|
||||
ocr_dict_ch: /data/project/KVM-privacy/deps/KVM/models/ocr/ch/ppocr_keys.txt
|
||||
ocr_dict_en: /data/project/KVM-privacy/deps/KVM/models/ocr/en/ppocr_keys.txt
|
||||
# MediaPipe face detection
|
||||
face_det: /data/project/KVM-privacy/deps/info-privacy-rs/deps/mediapipe-rknn/models/face_detection_short_range_rk3588.rknn
|
||||
|
||||
# NPU core assignment (RK3588 has 3 NPU cores)
|
||||
# OCR uses Core 0+1, Face uses Core 0 (serial with OCR det)
|
||||
# Embedding uses Core 2 (managed by mem-bridge, not this daemon)
|
||||
cores:
|
||||
ocr_det: core0
|
||||
ocr_rec: core1
|
||||
face_det: core0
|
||||
face_det: core0 # shares with det, runs serially via scheduler
|
||||
|
||||
scheduler:
|
||||
max_concurrent: 2
|
||||
queue_size: 32
|
||||
max_concurrent: 2 # max parallel NPU tasks
|
||||
queue_size: 32 # per-priority queue depth
|
||||
|
||||
+2
-2
@@ -7,13 +7,14 @@ Wants=network.target
|
||||
Type=simple
|
||||
User=pi
|
||||
Group=pi
|
||||
ExecStart=/usr/local/bin/npu-daemon --config /etc/npu-daemon/config.yaml
|
||||
ExecStart=/usr/sbin/npu-daemon --config /etc/npu-daemon/config.yaml
|
||||
|
||||
# Environment
|
||||
Environment=RUST_LOG=npu_daemon=info
|
||||
Environment=RKNN_LIB_DIR=/usr/lib
|
||||
|
||||
# CPU affinity: A55 cores (4-7) for NPU management overhead
|
||||
# NPU cores handle actual inference, CPU just coordinates
|
||||
CPUAffinity=4 5 6 7
|
||||
|
||||
# Resource limits
|
||||
@@ -30,7 +31,6 @@ StartLimitIntervalSec=60
|
||||
NoNewPrivileges=true
|
||||
ProtectSystem=strict
|
||||
ProtectHome=read-only
|
||||
ReadOnlyPaths=/usr/share/kvm-npu
|
||||
ReadWritePaths=/tmp
|
||||
PrivateTmp=true
|
||||
|
||||
|
||||
Vendored
+2
@@ -0,0 +1,2 @@
|
||||
/etc/kvm-privacy/pii_rules.yaml
|
||||
/etc/kvm-privacy/surnames.txt
|
||||
+9
@@ -0,0 +1,9 @@
|
||||
#!/bin/bash
|
||||
set -e
|
||||
case "$1" in
|
||||
purge)
|
||||
rm -rf /usr/lib/kvm-privacy
|
||||
rm -rf /etc/kvm-privacy
|
||||
;;
|
||||
esac
|
||||
exit 0
|
||||
+2
-1
@@ -43,6 +43,7 @@ confidential_keywords:
|
||||
#
|
||||
# det_model / rec_model / face.model 为空时,对应功能自动禁用。
|
||||
rknn:
|
||||
npu_daemon_url: "http://localhost:8004"
|
||||
ocr:
|
||||
det_model: /data/project/KVM-privacy/deps/info-privacy-rs/deps/ocr-rknn/models/PP-OCRv4/det/ppocrv4_det_ch_int8.rknn
|
||||
rec_model: /data/project/KVM-privacy/deps/info-privacy-rs/deps/ocr-rknn/models/PP-OCRv4/rec/ppocrv4_rec_ch_fp16.rknn
|
||||
@@ -60,7 +61,7 @@ rknn:
|
||||
ner:
|
||||
name:
|
||||
security_level: medium
|
||||
surnames_file: configs/surnames.txt
|
||||
surnames_file: /etc/kvm-privacy/surnames.txt
|
||||
address:
|
||||
security_level: medium
|
||||
triggers:
|
||||
|
||||
@@ -9,8 +9,9 @@ StartLimitBurst=5
|
||||
Type=simple
|
||||
User=pi
|
||||
Group=pi
|
||||
WorkingDirectory=/usr/lib/kvm-privacy
|
||||
ExecStart=/usr/local/bin/info-privacy-rs
|
||||
ExecStart=/usr/sbin/info-privacy-rs
|
||||
RuntimeDirectory=kvm-privacy
|
||||
RuntimeDirectoryMode=0750
|
||||
Environment=PORT=8001
|
||||
Environment=NPU_DAEMON_URL=http://localhost:8004
|
||||
Environment=PII_RULES=/etc/kvm-privacy/pii_rules.yaml
|
||||
|
||||
@@ -1,187 +0,0 @@
|
||||
#!/usr/bin/env python3
|
||||
"""Face Detection Worker Daemon — 供 Rust 子进程调用(JSON line 协议)
|
||||
|
||||
协议:
|
||||
stdin 每行一个 JSON:{"image_path": "/tmp/foo.jpg", "page": 1, "target": "rk3588"}
|
||||
stdout 每行一个 JSON:{"ok": true, "faces": [...]} 或 {"ok": false, "error": "..."}
|
||||
|
||||
每个 face:
|
||||
{"bbox": [x1,y1,x2,y2], "page": int, "score": float}
|
||||
|
||||
后端优先级:
|
||||
1. NPU Daemon HTTP API(NPU_DAEMON_URL 环境变量)
|
||||
2. mediapipe-rknn 板端后端(face_detection.FaceDetector)
|
||||
3. mediapipe_rknn Python 包(x86 开发机)
|
||||
"""
|
||||
import sys
|
||||
import os
|
||||
import json
|
||||
|
||||
_NPU_DAEMON_URL = os.environ.get("NPU_DAEMON_URL", "") # e.g. http://localhost:8004
|
||||
_MEDIAPIPE_SRC = os.environ.get("MEDIAPIPE_SRC", "/data/rockchip/mediapipe/src")
|
||||
_FACE_MODEL = os.environ.get("FACE_MODEL",
|
||||
"/data/rockchip/mediapipe/models/rknn/face_detection_short_range_rk3588.rknn")
|
||||
_TARGET = os.environ.get("RKNN_TARGET", "rk3588")
|
||||
|
||||
if _MEDIAPIPE_SRC not in sys.path:
|
||||
sys.path.insert(0, _MEDIAPIPE_SRC)
|
||||
|
||||
|
||||
# ── 后端 C:NPU Daemon HTTP API(推荐,集中调度)──────────────────
|
||||
def _run_npu_daemon(image_path: str, page: int):
|
||||
"""通过 NPU Daemon HTTP API 检测人脸。"""
|
||||
import io
|
||||
import requests
|
||||
from PIL import Image, ImageOps
|
||||
|
||||
# Fix EXIF orientation and get corrected dimensions
|
||||
img = Image.open(image_path)
|
||||
img = ImageOps.exif_transpose(img)
|
||||
img_w, img_h = img.size
|
||||
|
||||
# Re-encode with correct orientation for NPU Daemon
|
||||
buf = io.BytesIO()
|
||||
img.save(buf, format="JPEG", quality=90)
|
||||
buf.seek(0)
|
||||
|
||||
url = f"{_NPU_DAEMON_URL}/api/v1/face/detect"
|
||||
resp = requests.post(
|
||||
url,
|
||||
files={"image": ("image.jpg", buf, "image/jpeg")},
|
||||
timeout=10,
|
||||
)
|
||||
resp.raise_for_status()
|
||||
data = resp.json()
|
||||
|
||||
# Convert NPU Daemon normalized {x,y,width,height} → pixel [x1,y1,x2,y2]
|
||||
faces = []
|
||||
for face in data.get("faces", []):
|
||||
x1 = face["x"] * img_w
|
||||
y1 = face["y"] * img_h
|
||||
x2 = (face["x"] + face["width"]) * img_w
|
||||
y2 = (face["y"] + face["height"]) * img_h
|
||||
faces.append({
|
||||
"bbox": [float(x1), float(y1), float(x2), float(y2)],
|
||||
"page": page,
|
||||
"score": float(face["score"]),
|
||||
})
|
||||
return faces
|
||||
|
||||
|
||||
# ── 后端 A:mediapipe-rknn 项目(板端,face_detection.FaceDetector)────────
|
||||
def _try_board_backend(model_path: str, target: str):
|
||||
"""尝试板端 mediapipe-rknn 风格后端,失败返回 None。"""
|
||||
try:
|
||||
from face_detection import FaceDetector
|
||||
det = FaceDetector(model_path, target=target)
|
||||
return ("board", det)
|
||||
except Exception:
|
||||
return None
|
||||
|
||||
|
||||
def _run_board(det, image_path: str, page: int):
|
||||
"""使用板端 FaceDetector 检测,返回 blocks。"""
|
||||
import cv2
|
||||
img = cv2.imread(image_path)
|
||||
if img is None:
|
||||
raise ValueError(f"无法读取图像: {image_path}")
|
||||
|
||||
boxes, _kps, scores = det.detect(img)
|
||||
faces = []
|
||||
for bbox, score in zip(boxes, scores):
|
||||
if float(score) < 0.5:
|
||||
continue
|
||||
x1, y1, x2, y2 = float(bbox[0]), float(bbox[1]), float(bbox[2]), float(bbox[3])
|
||||
faces.append({
|
||||
"bbox": [x1, y1, x2, y2],
|
||||
"page": page,
|
||||
"score": float(score),
|
||||
})
|
||||
return faces
|
||||
|
||||
|
||||
# ── 后端 B:mediapipe_rknn Python 包(x86 开发机)───────────────────────────
|
||||
def _try_x86_backend(model_path: str):
|
||||
"""尝试 x86 mediapipe_rknn 包风格后端,失败返回 None。"""
|
||||
try:
|
||||
from mediapipe_rknn.solutions import FaceDetection
|
||||
det = FaceDetection(model_path=model_path)
|
||||
det.load()
|
||||
return ("x86", det)
|
||||
except Exception:
|
||||
return None
|
||||
|
||||
|
||||
def _run_x86(det, image_path: str, page: int):
|
||||
"""使用 x86 mediapipe_rknn FaceDetection 检测,返回 blocks。"""
|
||||
import cv2
|
||||
img = cv2.imread(image_path)
|
||||
if img is None:
|
||||
raise ValueError(f"无法读取图像: {image_path}")
|
||||
|
||||
results = det.detect(img)
|
||||
faces = []
|
||||
if results:
|
||||
for r in results:
|
||||
bbox = r.bbox # [x1, y1, x2, y2]
|
||||
score = getattr(r, "score", 1.0)
|
||||
if float(score) < 0.5:
|
||||
continue
|
||||
faces.append({
|
||||
"bbox": [float(v) for v in bbox],
|
||||
"page": page,
|
||||
"score": float(score),
|
||||
})
|
||||
return faces
|
||||
|
||||
|
||||
# ── 主循环 ────────────────────────────────────────────────────────────────────
|
||||
def main():
|
||||
sys.stdout.write(json.dumps({"ready": True}) + "\n")
|
||||
sys.stdout.flush()
|
||||
|
||||
use_npu_daemon = bool(_NPU_DAEMON_URL)
|
||||
if use_npu_daemon:
|
||||
print(f"[face_worker] Using NPU Daemon at {_NPU_DAEMON_URL}", file=sys.stderr)
|
||||
|
||||
backend_tag = None
|
||||
detector = None
|
||||
|
||||
for raw_line in sys.stdin:
|
||||
raw_line = raw_line.strip()
|
||||
if not raw_line:
|
||||
continue
|
||||
try:
|
||||
req = json.loads(raw_line)
|
||||
image_path = req["image_path"]
|
||||
page = int(req.get("page", 1))
|
||||
target = req.get("target", _TARGET)
|
||||
|
||||
if use_npu_daemon:
|
||||
faces = _run_npu_daemon(image_path, page)
|
||||
else:
|
||||
# 延迟初始化(只加载一次)
|
||||
if detector is None:
|
||||
result = _try_board_backend(_FACE_MODEL, target)
|
||||
if result:
|
||||
backend_tag, detector = result
|
||||
else:
|
||||
result = _try_x86_backend(_FACE_MODEL)
|
||||
if result:
|
||||
backend_tag, detector = result
|
||||
else:
|
||||
raise RuntimeError("无法初始化任何人脸检测后端(board 和 x86 均失败)")
|
||||
|
||||
if backend_tag == "board":
|
||||
faces = _run_board(detector, image_path, page)
|
||||
else:
|
||||
faces = _run_x86(detector, image_path, page)
|
||||
|
||||
sys.stdout.write(json.dumps({"ok": True, "faces": faces}) + "\n")
|
||||
except Exception as e:
|
||||
sys.stdout.write(json.dumps({"ok": False, "error": str(e)}) + "\n")
|
||||
sys.stdout.flush()
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -1,256 +0,0 @@
|
||||
#!/usr/bin/env python3
|
||||
"""OCR Worker Daemon — 供 Rust 子进程调用(JSON line 协议)
|
||||
|
||||
协议:
|
||||
stdin 每行一个 JSON:{"image_path": "/tmp/foo.jpg", "page": 1, "target": "rk3588"}
|
||||
stdout 每行一个 JSON:{"ok": true, "blocks": [...]} 或 {"ok": false, "error": "..."}
|
||||
|
||||
每个 block:
|
||||
{"text": str, "bbox": [x1,y1,x2,y2], "page": int, "layer": "image", "confidence": float}
|
||||
|
||||
后端优先级:
|
||||
1. NPU Daemon HTTP API(NPU_DAEMON_URL 环境变量)
|
||||
2. ppocr_rknn.PPOcrRknn(rknnlite,板端友好)
|
||||
3. ppocr_det + ppocr_rec(rknn full API)
|
||||
"""
|
||||
import sys
|
||||
import os
|
||||
import json
|
||||
|
||||
# ── 路径配置(均可通过环境变量覆盖)────────────────────────────────
|
||||
_NPU_DAEMON_URL = os.environ.get("NPU_DAEMON_URL", "") # e.g. http://localhost:8004
|
||||
_OCR_RKNN_DIR = os.environ.get("OCR_RKNN_DIR", "") # ppocr_rknn.py 所在目录
|
||||
_PPOCR_DIR = os.environ.get("OCR_PPOCR_DIR",
|
||||
"/data/rockchip/rknn_model_zoo/examples/PPOCR/PPOCR-System/python")
|
||||
_DET_MODEL = os.environ.get("OCR_DET_MODEL",
|
||||
"/data/rockchip/paddle_ocr/models/rknn/PP-OCRv4/det/ppocrv4_det_ch_int8.rknn")
|
||||
_REC_MODEL = os.environ.get("OCR_REC_MODEL",
|
||||
"/data/rockchip/paddle_ocr/models/rknn/PP-OCRv4/rec/ppocrv4_rec_ch_fp16.rknn")
|
||||
_DICT_PATH = os.environ.get("OCR_DICT_PATH",
|
||||
"/data/rockchip/rknn_model_zoo/examples/PPOCR/PPOCR-System/model/ppocr_keys_v1.txt")
|
||||
_TARGET = os.environ.get("RKNN_TARGET", "rk3588")
|
||||
|
||||
|
||||
# ── 图像预处理(EXIF 方向 + 缩放)───────────────────────────────
|
||||
_MAX_OCR_EDGE = int(os.environ.get("OCR_MAX_EDGE", "960"))
|
||||
|
||||
def _preprocess_image(image_path: str):
|
||||
"""Apply EXIF orientation and resize large images for OCR accuracy.
|
||||
|
||||
Returns (bytes, content_type) ready for HTTP upload.
|
||||
PP-OCRv4 det model input is 480×480; images much larger than ~960px
|
||||
lose text detail after downsampling. EXIF rotation is critical —
|
||||
the RKNN image crate does not auto-rotate.
|
||||
"""
|
||||
from PIL import Image, ImageOps
|
||||
import io
|
||||
|
||||
img = Image.open(image_path)
|
||||
|
||||
# Fix EXIF orientation (camera photos are often rotated)
|
||||
img = ImageOps.exif_transpose(img)
|
||||
|
||||
# Resize if too large — maintain aspect ratio
|
||||
max_edge = max(img.size)
|
||||
if max_edge > _MAX_OCR_EDGE:
|
||||
ratio = _MAX_OCR_EDGE / max_edge
|
||||
new_size = (int(img.size[0] * ratio), int(img.size[1] * ratio))
|
||||
img = img.resize(new_size, Image.LANCZOS)
|
||||
|
||||
buf = io.BytesIO()
|
||||
img.save(buf, format="JPEG", quality=90)
|
||||
buf.seek(0)
|
||||
return buf
|
||||
|
||||
|
||||
# ── 后端 C:NPU Daemon HTTP API(推荐,集中调度)──────────────────
|
||||
def _run_npu_daemon(image_path: str, page: int):
|
||||
"""通过 NPU Daemon HTTP API 运行 OCR。"""
|
||||
import requests
|
||||
|
||||
url = f"{_NPU_DAEMON_URL}/api/v1/ocr/analyze"
|
||||
image_buf = _preprocess_image(image_path)
|
||||
resp = requests.post(
|
||||
url,
|
||||
files={"image": ("image.jpg", image_buf, "image/jpeg")},
|
||||
data={"priority": "p3"},
|
||||
timeout=15,
|
||||
)
|
||||
resp.raise_for_status()
|
||||
data = resp.json()
|
||||
|
||||
# Convert NPU Daemon format {x,y,w,h} → worker format [x1,y1,x2,y2]
|
||||
blocks = []
|
||||
for region in data.get("regions", []):
|
||||
x, y = region["x"], region["y"]
|
||||
w, h = region["w"], region["h"]
|
||||
blocks.append({
|
||||
"text": region["text"],
|
||||
"bbox": [float(x), float(y), float(x + w), float(y + h)],
|
||||
"page": page,
|
||||
"layer": "image",
|
||||
"confidence": float(region["confidence"]),
|
||||
})
|
||||
return blocks
|
||||
|
||||
|
||||
# ── 后端 A:ppocr_rknn.PPOcrRknn(支持 rknnlite,推荐板端)──────────
|
||||
def _try_ppocr_rknn_backend(target: str):
|
||||
"""尝试使用 PPOcrRknn 后端,失败返回 None。"""
|
||||
if not _OCR_RKNN_DIR or not os.path.isdir(_OCR_RKNN_DIR):
|
||||
return None
|
||||
if _OCR_RKNN_DIR not in sys.path:
|
||||
sys.path.insert(0, _OCR_RKNN_DIR)
|
||||
try:
|
||||
from ppocr_rknn import PPOcrRknn
|
||||
# PPOcrRknn.__init__ 会 print() 初始化信息到 stdout,必须临时重定向到 stderr
|
||||
# 避免污染 Rust worker 的 JSON line 协议
|
||||
_orig_stdout = sys.stdout
|
||||
sys.stdout = sys.stderr
|
||||
try:
|
||||
ocr = PPOcrRknn(
|
||||
lang="ch",
|
||||
target=target,
|
||||
det_model_path=_DET_MODEL,
|
||||
rec_model_path=_REC_MODEL,
|
||||
dict_path=_DICT_PATH,
|
||||
use_cls=False,
|
||||
)
|
||||
finally:
|
||||
sys.stdout = _orig_stdout
|
||||
return ocr
|
||||
except Exception:
|
||||
return None
|
||||
|
||||
|
||||
def _run_ppocr_rknn(ocr, image_path: str, page: int):
|
||||
"""使用 PPOcrRknn 运行 OCR,返回 blocks 列表。"""
|
||||
results = ocr.run(image_path) # [(box[4,2], text, score), ...]
|
||||
blocks = []
|
||||
for box, text, score in results:
|
||||
if float(score) < 0.5:
|
||||
continue
|
||||
xs = box[:, 0]
|
||||
ys = box[:, 1]
|
||||
blocks.append({
|
||||
"text": str(text),
|
||||
"bbox": [float(xs.min()), float(ys.min()),
|
||||
float(xs.max()), float(ys.max())],
|
||||
"page": page,
|
||||
"layer": "image",
|
||||
"confidence": float(score),
|
||||
})
|
||||
return blocks
|
||||
|
||||
|
||||
# ── 后端 B:ppocr_det + ppocr_rec(需要 rknn full API)──────────────
|
||||
def _load_legacy_models(target: str):
|
||||
if _PPOCR_DIR not in sys.path:
|
||||
sys.path.insert(0, _PPOCR_DIR)
|
||||
import ppocr_det as predict_det
|
||||
import ppocr_rec as predict_rec
|
||||
|
||||
class Args:
|
||||
det_model_path = _DET_MODEL
|
||||
rec_model_path = _REC_MODEL
|
||||
dict_path = _DICT_PATH
|
||||
|
||||
args = Args()
|
||||
args.target = target
|
||||
args.device_id = None
|
||||
|
||||
detector = predict_det.TextDetector(args)
|
||||
recognizer = predict_rec.TextRecognizer(args)
|
||||
return detector, recognizer, predict_det
|
||||
|
||||
|
||||
def _run_legacy(detector, recognizer, predict_det, image_path: str, page: int):
|
||||
import cv2
|
||||
|
||||
img = cv2.imread(image_path)
|
||||
if img is None:
|
||||
raise ValueError(f"无法读取图像: {image_path}")
|
||||
|
||||
ori_im = img.copy()
|
||||
dt_boxes = detector.run(img)
|
||||
if dt_boxes is None or len(dt_boxes) == 0:
|
||||
return []
|
||||
|
||||
img_crop_list = [
|
||||
predict_det.get_rotate_crop_image(ori_im, box)
|
||||
for box in sorted(dt_boxes, key=lambda b: (b[0][1], b[0][0]))
|
||||
]
|
||||
rec_res = recognizer.run(img_crop_list)
|
||||
|
||||
blocks = []
|
||||
for box, rec in zip(dt_boxes, rec_res):
|
||||
if isinstance(rec, (list, tuple)) and rec:
|
||||
text = rec[0][0] if isinstance(rec[0], (list, tuple)) else rec[0]
|
||||
score = rec[0][1] if isinstance(rec[0], (list, tuple)) else (rec[1] if len(rec) > 1 else 1.0)
|
||||
else:
|
||||
continue
|
||||
if float(score) < 0.5:
|
||||
continue
|
||||
xs, ys = box[:, 0], box[:, 1]
|
||||
blocks.append({
|
||||
"text": str(text),
|
||||
"bbox": [float(xs.min()), float(ys.min()),
|
||||
float(xs.max()), float(ys.max())],
|
||||
"page": page,
|
||||
"layer": "image",
|
||||
"confidence": float(score),
|
||||
})
|
||||
return blocks
|
||||
|
||||
|
||||
# ── 主循环 ────────────────────────────────────────────────────────────
|
||||
def main():
|
||||
sys.stdout.write(json.dumps({"ready": True}) + "\n")
|
||||
sys.stdout.flush()
|
||||
|
||||
use_npu_daemon = bool(_NPU_DAEMON_URL)
|
||||
if use_npu_daemon:
|
||||
print(f"[ocr_worker] Using NPU Daemon at {_NPU_DAEMON_URL}", file=sys.stderr)
|
||||
|
||||
# 延迟初始化(仅 local RKNN 后端需要)
|
||||
ppocr_rknn = None # 后端 A
|
||||
det = rec = det_mod = None # 后端 B
|
||||
backend = None
|
||||
target = _TARGET
|
||||
|
||||
for raw_line in sys.stdin:
|
||||
raw_line = raw_line.strip()
|
||||
if not raw_line:
|
||||
continue
|
||||
try:
|
||||
req = json.loads(raw_line)
|
||||
image_path = req["image_path"]
|
||||
page = int(req.get("page", 1))
|
||||
req_target = req.get("target", target)
|
||||
|
||||
if use_npu_daemon:
|
||||
blocks = _run_npu_daemon(image_path, page)
|
||||
else:
|
||||
# 首次或 target 变更时加载模型
|
||||
if backend is None or req_target != target:
|
||||
target = req_target
|
||||
ppocr_rknn = _try_ppocr_rknn_backend(target)
|
||||
if ppocr_rknn is not None:
|
||||
backend = "ppocr_rknn"
|
||||
else:
|
||||
det, rec, det_mod = _load_legacy_models(target)
|
||||
backend = "legacy"
|
||||
|
||||
if backend == "ppocr_rknn":
|
||||
blocks = _run_ppocr_rknn(ppocr_rknn, image_path, page)
|
||||
else:
|
||||
blocks = _run_legacy(det, rec, det_mod, image_path, page)
|
||||
|
||||
sys.stdout.write(json.dumps({"ok": True, "blocks": blocks}) + "\n")
|
||||
except Exception as e:
|
||||
sys.stdout.write(json.dumps({"ok": False, "error": str(e)}) + "\n")
|
||||
sys.stdout.flush()
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
Vendored
+3
-4
@@ -1,10 +1,9 @@
|
||||
Package: kvm-rkllm
|
||||
Version: 1.0.0-1
|
||||
Version: 2.0.0-1
|
||||
Architecture: arm64
|
||||
Maintainer: KVM-Privacy <noreply@kvm-privacy.local>
|
||||
Depends: python3 (>= 3.9)
|
||||
Recommends: python3-fastapi, python3-uvicorn, python3-pydantic
|
||||
Description: RKLLM NPU inference server (Qwen on RK3588)
|
||||
Depends: libc6
|
||||
Description: RKLLM NPU inference server (Rust, Qwen on RK3588)
|
||||
OpenAI-compatible chat API backed by RKLLM NPU runtime.
|
||||
Includes librkllmrt.so for RK3588 NPU acceleration.
|
||||
LLM model not included - set RKLLM_MODEL env var.
|
||||
|
||||
Vendored
+17
-1
@@ -5,7 +5,23 @@ if [ "$1" = "configure" ]; then
|
||||
# Ensure librkllmrt.so is discoverable by the dynamic linker
|
||||
ldconfig 2>/dev/null || true
|
||||
|
||||
# Create model directory and link known model locations
|
||||
mkdir -p /var/lib/kvm-rkllm/models
|
||||
chown pi:pi /var/lib/kvm-rkllm /var/lib/kvm-rkllm/models
|
||||
|
||||
# Auto-link model from known locations if not already present
|
||||
MODEL_NAME="Qwen3-0.6B-rk3588-w8a8-opt-1-hybrid-ratio-0.5.rkllm"
|
||||
if [ ! -e "/var/lib/kvm-rkllm/models/$MODEL_NAME" ]; then
|
||||
for dir in /opt/fileguard/models /home/pi/models; do
|
||||
if [ -f "$dir/$MODEL_NAME" ]; then
|
||||
ln -sf "$dir/$MODEL_NAME" "/var/lib/kvm-rkllm/models/$MODEL_NAME"
|
||||
echo "kvm-rkllm: linked model from $dir"
|
||||
break
|
||||
fi
|
||||
done
|
||||
fi
|
||||
|
||||
systemctl daemon-reload
|
||||
systemctl enable rkllm-server.service
|
||||
systemctl start rkllm-server.service || true
|
||||
systemctl start rkllm-server.service 2>/dev/null || true
|
||||
fi
|
||||
|
||||
+8
@@ -0,0 +1,8 @@
|
||||
#!/bin/bash
|
||||
set -e
|
||||
case "$1" in
|
||||
purge)
|
||||
rm -rf /var/lib/kvm-rkllm
|
||||
;;
|
||||
esac
|
||||
exit 0
|
||||
Vendored
-3
@@ -1,3 +0,0 @@
|
||||
#!/bin/bash
|
||||
export PYTHONPATH=/usr/lib/kvm-rkllm
|
||||
exec python3 -m rkllm_server "$@"
|
||||
@@ -9,8 +9,9 @@ StartLimitBurst=5
|
||||
Type=simple
|
||||
User=pi
|
||||
Group=pi
|
||||
WorkingDirectory=/usr/lib/kvm-privacy
|
||||
ExecStart=/usr/local/bin/info-privacy-rs
|
||||
ExecStart=/usr/sbin/info-privacy-rs
|
||||
RuntimeDirectory=kvm-privacy
|
||||
RuntimeDirectoryMode=0750
|
||||
Environment=PORT=8001
|
||||
Environment=NPU_DAEMON_URL=http://localhost:8004
|
||||
Environment=PII_RULES=/etc/kvm-privacy/pii_rules.yaml
|
||||
|
||||
@@ -1,5 +1,5 @@
|
||||
[Unit]
|
||||
Description=Mem-Bridge Memory Service
|
||||
Description=Embed-DB Memory + Router Service (Rust, ports 8001+8002)
|
||||
After=network.target
|
||||
Wants=network.target
|
||||
StartLimitInterval=300
|
||||
@@ -9,16 +9,18 @@ StartLimitBurst=5
|
||||
Type=simple
|
||||
User=pi
|
||||
Group=pi
|
||||
WorkingDirectory=/home/pi/Desktop/embed-db
|
||||
Environment=PYTHONPATH=/home/pi/Desktop/embed-db/src
|
||||
EnvironmentFile=-/home/pi/Desktop/embed-db/.env
|
||||
ExecStart=/home/pi/Desktop/embed-db/venv/bin/python server.py --service memory
|
||||
EnvironmentFile=-/etc/kvm-bridge/bridge.env
|
||||
Environment=MEMORY_PORT=8003
|
||||
Environment=ROUTER_PORT=8002
|
||||
ExecStart=/usr/sbin/embed-db
|
||||
Restart=on-failure
|
||||
RestartSec=10
|
||||
MemoryMax=512M
|
||||
# RK3588: pin to A55 small cores 2-3 (Python FAISS + embedding)
|
||||
MemoryMax=1G
|
||||
CPUAffinity=2 3
|
||||
LimitNOFILE=4096
|
||||
StateDirectory=kvm-bridge
|
||||
RuntimeDirectory=kvm-bridge
|
||||
RuntimeDirectoryMode=0750
|
||||
StandardOutput=journal
|
||||
StandardError=journal
|
||||
|
||||
|
||||
@@ -1,27 +0,0 @@
|
||||
[Unit]
|
||||
Description=Mem-Bridge Router Service
|
||||
After=network.target mem-bridge-memory.service
|
||||
Wants=network.target
|
||||
Requires=mem-bridge-memory.service
|
||||
StartLimitInterval=300
|
||||
StartLimitBurst=5
|
||||
|
||||
[Service]
|
||||
Type=simple
|
||||
User=pi
|
||||
Group=pi
|
||||
WorkingDirectory=/home/pi/Desktop/embed-db
|
||||
Environment=PYTHONPATH=/home/pi/Desktop/embed-db/src
|
||||
EnvironmentFile=-/home/pi/Desktop/embed-db/.env
|
||||
ExecStart=/home/pi/Desktop/embed-db/venv/bin/python server.py --service router
|
||||
Restart=on-failure
|
||||
RestartSec=10
|
||||
MemoryMax=2G
|
||||
# RK3588: pin to A55 small cores 2-3 (AI routing + LLM orchestration)
|
||||
CPUAffinity=2 3
|
||||
LimitNOFILE=4096
|
||||
StandardOutput=journal
|
||||
StandardError=journal
|
||||
|
||||
[Install]
|
||||
WantedBy=multi-user.target
|
||||
@@ -1,18 +0,0 @@
|
||||
[Unit]
|
||||
Description=KVM Privacy Gateway REST API
|
||||
After=network.target privacy-gateway.service
|
||||
Wants=privacy-gateway.service
|
||||
|
||||
[Service]
|
||||
Type=simple
|
||||
User=root
|
||||
WorkingDirectory=/data/project/KVM-privacy/services
|
||||
ExecStart=/usr/bin/python3 -m privacy_gateway
|
||||
Restart=on-failure
|
||||
RestartSec=5
|
||||
MemoryMax=128M
|
||||
StandardOutput=journal
|
||||
StandardError=journal
|
||||
|
||||
[Install]
|
||||
WantedBy=multi-user.target
|
||||
@@ -1,26 +0,0 @@
|
||||
[Unit]
|
||||
Description=KVM Privacy Gateway (mitmproxy)
|
||||
After=network.target info-privacy.service
|
||||
Wants=network.target info-privacy.service
|
||||
StartLimitInterval=300
|
||||
StartLimitBurst=5
|
||||
|
||||
[Service]
|
||||
Type=simple
|
||||
User=root
|
||||
WorkingDirectory=/data/project/KVM-privacy/services/privacy_gateway
|
||||
ExecStart=/usr/local/bin/mitmdump \
|
||||
--mode transparent \
|
||||
--listen-host 0.0.0.0 \
|
||||
--listen-port 8888 \
|
||||
--ssl-insecure \
|
||||
-s /data/project/KVM-privacy/services/privacy_gateway/addon.py
|
||||
Restart=on-failure
|
||||
RestartSec=10
|
||||
MemoryMax=768M
|
||||
LimitNOFILE=8192
|
||||
StandardOutput=journal
|
||||
StandardError=journal
|
||||
|
||||
[Install]
|
||||
WantedBy=multi-user.target
|
||||
@@ -1,17 +1,16 @@
|
||||
[Unit]
|
||||
Description=RKLLM Server (Qwen3-0.6B on RK3588 NPU, port 8891)
|
||||
Description=RKLLM Server (Rust, Qwen3-0.6B on RK3588 NPU, port 8891)
|
||||
Documentation=https://github.com/airockchip/rknn-llm
|
||||
After=network.target
|
||||
Before=mem-bridge-router.service
|
||||
Before=embed-db.service
|
||||
|
||||
[Service]
|
||||
Type=simple
|
||||
User=pi
|
||||
WorkingDirectory=/usr/lib/kvm-rkllm
|
||||
Environment=RKLLM_MODEL=/opt/fileguard/models/Qwen3-0.6B-rk3588-w8a8-opt-1-hybrid-ratio-0.5.rkllm
|
||||
Environment=RKLLM_MODEL=/var/lib/kvm-rkllm/models/Qwen3-0.6B-rk3588-w8a8-opt-1-hybrid-ratio-0.5.rkllm
|
||||
Environment=RKLLM_LIB=/usr/lib/kvm-rkllm/librkllmrt.so
|
||||
Environment=RKLLM_PORT=8891
|
||||
Environment=PYTHONPATH=/usr/lib/kvm-rkllm
|
||||
ExecStart=/usr/bin/python3 -m rkllm_server
|
||||
ExecStart=/usr/sbin/rkllm-server
|
||||
Restart=on-failure
|
||||
RestartSec=5
|
||||
CPUAffinity=0 1
|
||||
|
||||
@@ -1,21 +0,0 @@
|
||||
[Unit]
|
||||
Description=KVM-Privacy Workflow Dashboard (port 9099)
|
||||
After=network.target npu-daemon.service
|
||||
Wants=npu-daemon.service
|
||||
|
||||
[Service]
|
||||
Type=simple
|
||||
User=pi
|
||||
Group=pi
|
||||
WorkingDirectory=/data/project/KVM-privacy/tools/workflow-dashboard
|
||||
Environment=PYTHONUNBUFFERED=1
|
||||
ExecStart=/usr/bin/python3 -m workflow_dashboard
|
||||
Restart=on-failure
|
||||
RestartSec=5
|
||||
MemoryMax=256M
|
||||
CPUAffinity=4 5 6 7
|
||||
StandardOutput=journal
|
||||
StandardError=journal
|
||||
|
||||
[Install]
|
||||
WantedBy=multi-user.target
|
||||
@@ -1,3 +1,5 @@
|
||||
# Core services are deployed via DEB packages and systemd.
|
||||
# This file only runs the workflow dashboard for monitoring.
|
||||
services:
|
||||
workflow-dashboard:
|
||||
build:
|
||||
|
||||
@@ -216,35 +216,35 @@ echo "==> kvm-mitm deb: $MITM_DEB"
|
||||
# kvm-agent (Rust)
|
||||
AGENT_VER=$(grep "^Version:" debian/kvm-agent/DEBIAN/control | awk '{print $2}')
|
||||
AGENT_DEB="build/deb/kvm-agent_${AGENT_VER}_${ARCH}.deb"
|
||||
chmod 755 debian/kvm-agent/DEBIAN/postinst debian/kvm-agent/DEBIAN/prerm
|
||||
chmod 755 debian/kvm-agent/DEBIAN/postinst debian/kvm-agent/DEBIAN/prerm debian/kvm-agent/DEBIAN/postrm
|
||||
dpkg-deb --build debian/kvm-agent "$AGENT_DEB"
|
||||
echo "==> kvm-agent deb: $AGENT_DEB"
|
||||
|
||||
# kvm-bridge (Rust embed-db)
|
||||
BRIDGE_VER=$(grep "^Version:" debian/kvm-bridge/DEBIAN/control | awk '{print $2}')
|
||||
BRIDGE_DEB="build/deb/kvm-bridge_${BRIDGE_VER}_${ARCH}.deb"
|
||||
chmod 755 debian/kvm-bridge/DEBIAN/postinst debian/kvm-bridge/DEBIAN/prerm
|
||||
chmod 755 debian/kvm-bridge/DEBIAN/postinst debian/kvm-bridge/DEBIAN/prerm debian/kvm-bridge/DEBIAN/postrm
|
||||
dpkg-deb --build debian/kvm-bridge "$BRIDGE_DEB"
|
||||
echo "==> kvm-bridge deb: $BRIDGE_DEB"
|
||||
|
||||
# kvm-rkllm (Rust + librkllmrt.so)
|
||||
RKLLM_VER=$(grep "^Version:" debian/kvm-rkllm/DEBIAN/control | awk '{print $2}')
|
||||
RKLLM_DEB="build/deb/kvm-rkllm_${RKLLM_VER}_arm64.deb"
|
||||
chmod 755 debian/kvm-rkllm/DEBIAN/postinst debian/kvm-rkllm/DEBIAN/prerm
|
||||
chmod 755 debian/kvm-rkllm/DEBIAN/postinst debian/kvm-rkllm/DEBIAN/prerm debian/kvm-rkllm/DEBIAN/postrm
|
||||
dpkg-deb --build debian/kvm-rkllm "$RKLLM_DEB"
|
||||
echo "==> kvm-rkllm deb: $RKLLM_DEB"
|
||||
|
||||
# kvm-npu (Rust)
|
||||
NPU_VER=$(grep "^Version:" debian/kvm-npu/DEBIAN/control | awk '{print $2}')
|
||||
NPU_DEB="build/deb/kvm-npu_${NPU_VER}_arm64.deb"
|
||||
chmod 755 debian/kvm-npu/DEBIAN/postinst debian/kvm-npu/DEBIAN/prerm
|
||||
chmod 755 debian/kvm-npu/DEBIAN/postinst debian/kvm-npu/DEBIAN/prerm debian/kvm-npu/DEBIAN/postrm
|
||||
dpkg-deb --build debian/kvm-npu "$NPU_DEB"
|
||||
echo "==> kvm-npu deb: $NPU_DEB"
|
||||
|
||||
# kvm-privacy (Rust, no Python workers)
|
||||
PRIVACY_VER=$(grep "^Version:" debian/kvm-privacy/DEBIAN/control | awk '{print $2}')
|
||||
PRIVACY_DEB="build/deb/kvm-privacy_${PRIVACY_VER}_arm64.deb"
|
||||
chmod 755 debian/kvm-privacy/DEBIAN/postinst debian/kvm-privacy/DEBIAN/prerm
|
||||
chmod 755 debian/kvm-privacy/DEBIAN/postinst debian/kvm-privacy/DEBIAN/prerm debian/kvm-privacy/DEBIAN/postrm
|
||||
dpkg-deb --build debian/kvm-privacy "$PRIVACY_DEB"
|
||||
echo "==> kvm-privacy deb: $PRIVACY_DEB"
|
||||
|
||||
|
||||
+18
-27
@@ -1,38 +1,29 @@
|
||||
#!/usr/bin/env bash
|
||||
# setup_gateway.sh — 安装 Privacy Gateway (mitmproxy + CA 证书 + systemd)
|
||||
# setup_gateway.sh — Deploy Rust Privacy Gateway + related services via DEB packages
|
||||
set -euo pipefail
|
||||
|
||||
DEVICE_IP="${DEVICE_IP:-192.168.123.181}"
|
||||
DEVICE_USER="${DEVICE_USER:-pi}"
|
||||
SCRIPT_DIR="$(cd "$(dirname "${BASH_SOURCE[0]}")" && pwd)"
|
||||
DEPLOY_DIR="${SCRIPT_DIR}/../deploy/systemd"
|
||||
GATEWAY_DIR="${SCRIPT_DIR}/../services/privacy_gateway"
|
||||
BUILD_DIR="${SCRIPT_DIR}/../build/deb"
|
||||
|
||||
echo "=== [1/4] 安装 mitmproxy ==="
|
||||
ssh "${DEVICE_USER}@${DEVICE_IP}" "pip3 install mitmproxy httpx cryptography 2>&1 | tail -5"
|
||||
echo "=== [1/3] Build DEB packages ==="
|
||||
bash "${SCRIPT_DIR}/build-debs.sh"
|
||||
|
||||
echo "=== [2/4] 上传 Privacy Gateway 服务文件 ==="
|
||||
ssh "${DEVICE_USER}@${DEVICE_IP}" "mkdir -p /data/project/KVM-privacy/services/privacy_gateway"
|
||||
scp "${GATEWAY_DIR}"/*.py "${DEVICE_USER}@${DEVICE_IP}:/data/project/KVM-privacy/services/privacy_gateway/"
|
||||
scp "${GATEWAY_DIR}/ai_domains.txt" "${DEVICE_USER}@${DEVICE_IP}:/data/project/KVM-privacy/services/privacy_gateway/"
|
||||
|
||||
echo "=== [3/4] 生成 CA 证书 ==="
|
||||
ssh "${DEVICE_USER}@${DEVICE_IP}" "python3 -c '
|
||||
import sys; sys.path.insert(0, \"/data/project/KVM-privacy/services/privacy_gateway\")
|
||||
from cert_manager import ensure_ca
|
||||
cert, key = ensure_ca()
|
||||
print(f\"CA 证书: {cert}\")
|
||||
'"
|
||||
|
||||
echo "=== [4/4] 安装 systemd 服务 ==="
|
||||
for svc in info-privacy mem-bridge-memory mem-bridge-router privacy-gateway; do
|
||||
if [ -f "${DEPLOY_DIR}/${svc}.service" ]; then
|
||||
scp "${DEPLOY_DIR}/${svc}.service" "${DEVICE_USER}@${DEVICE_IP}:/tmp/${svc}.service"
|
||||
ssh "${DEVICE_USER}@${DEVICE_IP}" "echo '123123' | sudo -S cp /tmp/${svc}.service /etc/systemd/system/${svc}.service"
|
||||
echo ""
|
||||
echo "=== [2/3] Upload DEB packages ==="
|
||||
for deb in kvm-mitm kvm-privacy kvm-npu kvm-bridge kvm-rkllm; do
|
||||
DEB_FILE=$(ls "${BUILD_DIR}/${deb}_"*.deb 2>/dev/null | tail -1)
|
||||
if [ -n "$DEB_FILE" ]; then
|
||||
echo " Uploading $(basename "$DEB_FILE")"
|
||||
scp "$DEB_FILE" "${DEVICE_USER}@${DEVICE_IP}:/tmp/"
|
||||
fi
|
||||
done
|
||||
ssh "${DEVICE_USER}@${DEVICE_IP}" "echo '123123' | sudo -S systemctl daemon-reload"
|
||||
|
||||
echo "=== 完成!==="
|
||||
echo "启用服务: ssh ${DEVICE_USER}@${DEVICE_IP} 'sudo systemctl enable --now privacy-gateway'"
|
||||
echo "下载 CA 证书: curl http://${DEVICE_IP}:8080/api/v1/privacy/cert -H 'Authorization: Bearer <token>' -o kvm-ca.crt"
|
||||
echo ""
|
||||
echo "=== [3/3] Install on device ==="
|
||||
echo "Run on device:"
|
||||
echo " sudo dpkg -i /tmp/kvm-npu_*.deb /tmp/kvm-privacy_*.deb /tmp/kvm-mitm_*.deb /tmp/kvm-rkllm_*.deb /tmp/kvm-bridge_*.deb"
|
||||
echo ""
|
||||
echo "Verify:"
|
||||
echo " bash tools/smoke-test.sh"
|
||||
|
||||
@@ -219,27 +219,18 @@ verify_bridge() {
|
||||
return
|
||||
fi
|
||||
|
||||
# Service files
|
||||
for svc in mem-bridge-memory.service mem-bridge-router.service; do
|
||||
if systemctl list-unit-files "$svc" >/dev/null 2>&1; then
|
||||
pass "kvm-bridge $svc unit file present"
|
||||
else
|
||||
fail "kvm-bridge $svc unit file missing"
|
||||
fi
|
||||
done
|
||||
|
||||
# Python source installed
|
||||
if [ -d /usr/lib/kvm-bridge/mem_bridge ]; then
|
||||
pass "kvm-bridge mem_bridge source installed"
|
||||
# Service file (single embed-db service replaces mem-bridge-memory + router)
|
||||
if systemctl list-unit-files embed-db.service >/dev/null 2>&1; then
|
||||
pass "kvm-bridge embed-db.service unit file present"
|
||||
else
|
||||
fail "kvm-bridge mem_bridge source missing"
|
||||
fail "kvm-bridge embed-db.service unit file missing"
|
||||
fi
|
||||
|
||||
# Requirements file
|
||||
if [ -f /usr/lib/kvm-bridge/requirements.txt ]; then
|
||||
pass "kvm-bridge requirements.txt installed"
|
||||
# Rust binary installed
|
||||
if [ -x /usr/sbin/embed-db ]; then
|
||||
pass "kvm-bridge embed-db binary installed"
|
||||
else
|
||||
fail "kvm-bridge requirements.txt missing"
|
||||
fail "kvm-bridge embed-db binary missing"
|
||||
fi
|
||||
}
|
||||
|
||||
|
||||
@@ -114,10 +114,17 @@ impl ServiceConfig {
|
||||
}
|
||||
|
||||
fn default_memory_port() -> u16 {
|
||||
8001
|
||||
std::env::var("MEMORY_PORT")
|
||||
.ok()
|
||||
.and_then(|v| v.parse().ok())
|
||||
.unwrap_or(8003)
|
||||
}
|
||||
|
||||
fn default_router_port() -> u16 {
|
||||
8002
|
||||
std::env::var("ROUTER_PORT")
|
||||
.ok()
|
||||
.and_then(|v| v.parse().ok())
|
||||
.unwrap_or(8002)
|
||||
}
|
||||
fn default_db_dir() -> PathBuf {
|
||||
std::env::var("HOME")
|
||||
|
||||
@@ -27,6 +27,8 @@ pub mod backend;
|
||||
#[cfg(feature = "onnx")]
|
||||
pub mod onnx;
|
||||
pub mod pool;
|
||||
#[cfg(feature = "rknn")]
|
||||
pub mod rknn;
|
||||
pub mod tokenize;
|
||||
|
||||
use std::path::PathBuf;
|
||||
@@ -40,6 +42,8 @@ use tokenize::tokenize_batch;
|
||||
pub use backend::MockBackend;
|
||||
#[cfg(feature = "onnx")]
|
||||
pub use onnx::OnnxBackend;
|
||||
#[cfg(feature = "rknn")]
|
||||
pub use rknn::RknnBackend;
|
||||
|
||||
/// Errors returned by the embedder crate.
|
||||
#[derive(Debug, thiserror::Error)]
|
||||
|
||||
@@ -0,0 +1,37 @@
|
||||
//! RKNN NPU embedding backend (stub).
|
||||
//!
|
||||
//! This module will integrate with the NPU Daemon's `rknn` crate
|
||||
//! to run embedding inference on RK3588 NPU Core 2.
|
||||
|
||||
use crate::backend::{EmbedBackend, EmbedError};
|
||||
|
||||
/// RKNN NPU embedding backend.
|
||||
///
|
||||
/// Currently a placeholder — will be implemented when integrating
|
||||
/// with the `npu_daemon` RKNN runtime.
|
||||
pub struct RknnBackend {
|
||||
dimensions: usize,
|
||||
}
|
||||
|
||||
impl RknnBackend {
|
||||
/// Create a new RKNN backend (stub).
|
||||
pub fn new(dimensions: usize) -> Self {
|
||||
Self { dimensions }
|
||||
}
|
||||
}
|
||||
|
||||
impl EmbedBackend for RknnBackend {
|
||||
fn embed_batch(
|
||||
&self,
|
||||
_input_ids: &[Vec<i64>],
|
||||
_attention_mask: &[Vec<i64>],
|
||||
) -> Result<Vec<Vec<f32>>, EmbedError> {
|
||||
Err(EmbedError::Backend(
|
||||
"RKNN backend not yet implemented".into(),
|
||||
))
|
||||
}
|
||||
|
||||
fn dimensions(&self) -> usize {
|
||||
self.dimensions
|
||||
}
|
||||
}
|
||||
@@ -115,6 +115,30 @@ impl TurnStore {
|
||||
Ok(result)
|
||||
}
|
||||
|
||||
/// Retrieve turns by a list of FAISS IDs.
|
||||
pub fn get_by_faiss_ids(&self, faiss_ids: &[i64]) -> Result<Vec<TurnRow>, MetaStoreError> {
|
||||
if faiss_ids.is_empty() {
|
||||
return Ok(Vec::new());
|
||||
}
|
||||
let placeholders: Vec<String> = faiss_ids.iter().map(|_| "?".to_string()).collect();
|
||||
let sql = format!(
|
||||
"SELECT id, session_id, role, content, faiss_id, token_count, created_at
|
||||
FROM turns WHERE faiss_id IN ({})",
|
||||
placeholders.join(",")
|
||||
);
|
||||
let mut stmt = self.conn.prepare(&sql)?;
|
||||
let params: Vec<&dyn rusqlite::types::ToSql> = faiss_ids
|
||||
.iter()
|
||||
.map(|id| id as &dyn rusqlite::types::ToSql)
|
||||
.collect();
|
||||
let rows = stmt.query_map(params.as_slice(), row_to_turn)?;
|
||||
let mut result = Vec::new();
|
||||
for row in rows {
|
||||
result.push(row?);
|
||||
}
|
||||
Ok(result)
|
||||
}
|
||||
|
||||
/// Return the next available FAISS ID (max + 1, or 0 if empty).
|
||||
pub fn next_faiss_id(&self) -> Result<i64, MetaStoreError> {
|
||||
let val: Option<i64> = self.conn.query_row(
|
||||
|
||||
Generated
+747
-9
File diff suppressed because it is too large
Load Diff
@@ -18,6 +18,8 @@ members = [
|
||||
"crates/validator",
|
||||
"crates/llm-planner",
|
||||
"crates/hybrid-planner",
|
||||
"crates/runner",
|
||||
"crates/api-server",
|
||||
"crates/agent-core",
|
||||
]
|
||||
|
||||
@@ -69,8 +71,19 @@ window-manager = { path = "crates/window-manager" }
|
||||
validator = { path = "crates/validator" }
|
||||
llm-planner = { path = "crates/llm-planner" }
|
||||
hybrid-planner = { path = "crates/hybrid-planner" }
|
||||
runner = { path = "crates/runner" }
|
||||
api-server = { path = "crates/api-server" }
|
||||
agent-core = { path = "crates/agent-core" }
|
||||
|
||||
# database
|
||||
sqlx = { version = "0.8", features = ["runtime-tokio", "mysql", "json"] }
|
||||
|
||||
# streaming
|
||||
tokio-stream = { version = "0.1", features = ["sync"] }
|
||||
|
||||
# HTTP middleware
|
||||
tower-http = { version = "0.6", features = ["cors"] }
|
||||
|
||||
# base64
|
||||
base64 = "0.22"
|
||||
|
||||
|
||||
@@ -25,6 +25,8 @@ window-manager = { workspace = true }
|
||||
validator = { workspace = true }
|
||||
llm-planner = { workspace = true }
|
||||
hybrid-planner = { workspace = true }
|
||||
runner = { workspace = true }
|
||||
api-server = { workspace = true }
|
||||
serde = { workspace = true }
|
||||
serde_json = { workspace = true }
|
||||
tokio = { workspace = true }
|
||||
|
||||
@@ -130,6 +130,9 @@ impl KvmAgent {
|
||||
.run_task_loop(task_description, &mut action_history, &mut step_records)
|
||||
.await;
|
||||
|
||||
// Safe cleanup: dismiss dialogs, close stray windows
|
||||
self.safe_cleanup(3).await;
|
||||
|
||||
// Cleanup: release HID control + restore privacy mode
|
||||
tracing::info!("Releasing HID control");
|
||||
if let Err(e) = self.kvm.release_control().await {
|
||||
@@ -148,6 +151,9 @@ impl KvmAgent {
|
||||
action_history: &mut Vec<String>,
|
||||
step_records: &mut Vec<StepRecord>,
|
||||
) -> TaskResult {
|
||||
let mut consecutive_llm_errors: usize = 0;
|
||||
const MAX_CONSECUTIVE_LLM_ERRORS: usize = 3;
|
||||
|
||||
for step in 0..self.max_steps {
|
||||
if !self.running {
|
||||
return TaskResult {
|
||||
@@ -232,6 +238,34 @@ impl KvmAgent {
|
||||
action_history.push(action_desc.clone());
|
||||
tracing::info!(action = action_desc.as_str(), "Action");
|
||||
|
||||
// ── Consecutive LLM error check ────────────────
|
||||
if let Action::Wait { reason, .. } = &action {
|
||||
if reason.contains("LLM error") {
|
||||
consecutive_llm_errors += 1;
|
||||
if consecutive_llm_errors >= MAX_CONSECUTIVE_LLM_ERRORS {
|
||||
tracing::error!(
|
||||
count = consecutive_llm_errors,
|
||||
"LLM failed {} times consecutively — aborting task",
|
||||
consecutive_llm_errors,
|
||||
);
|
||||
return TaskResult {
|
||||
success: false,
|
||||
steps_taken: step + 1,
|
||||
final_reason: format!(
|
||||
"LLM unavailable after {} consecutive failures: {}",
|
||||
consecutive_llm_errors, reason,
|
||||
),
|
||||
actions: action_history.clone(),
|
||||
step_records: step_records.clone(),
|
||||
};
|
||||
}
|
||||
} else {
|
||||
consecutive_llm_errors = 0;
|
||||
}
|
||||
} else {
|
||||
consecutive_llm_errors = 0;
|
||||
}
|
||||
|
||||
// ── Done? ─────────────────────────────────────
|
||||
if action.action_type() == "done" {
|
||||
return TaskResult {
|
||||
@@ -634,6 +668,72 @@ impl KvmAgent {
|
||||
}
|
||||
}
|
||||
}
|
||||
/// Multi-round safe cleanup: dismiss shutdown dialogs, save prompts, close stray windows.
|
||||
///
|
||||
/// Each round: screenshot → OCR → detect state → act.
|
||||
/// Stops early when desktop is reached or max rounds exceeded.
|
||||
async fn safe_cleanup(&self, max_rounds: usize) {
|
||||
for round in 0..max_rounds {
|
||||
let screenshot = match self.kvm.screenshot().await {
|
||||
Ok(ss) => ss,
|
||||
Err(_) => return,
|
||||
};
|
||||
|
||||
let scene = match self.perceive_scene(&screenshot).await {
|
||||
Some(s) => s,
|
||||
None => return,
|
||||
};
|
||||
|
||||
let text_lower = scene.full_text.to_lowercase();
|
||||
|
||||
// Shutdown/restart dialog → click Cancel
|
||||
if text_lower.contains("shut down")
|
||||
|| text_lower.contains("关机")
|
||||
|| text_lower.contains("restart")
|
||||
|| text_lower.contains("重启")
|
||||
{
|
||||
tracing::info!(round = round + 1, "Cleanup: shutdown dialog detected, clicking Cancel");
|
||||
if mouse_ops::click_element(&self.kvm, "Cancel", &scene, 100).await
|
||||
|| mouse_ops::click_element(&self.kvm, "取消", &scene, 100).await
|
||||
{
|
||||
sleep(Duration::from_millis(500)).await;
|
||||
continue;
|
||||
}
|
||||
}
|
||||
|
||||
// Save prompt → click Don't Save / 不保存
|
||||
if text_lower.contains("save")
|
||||
|| text_lower.contains("保存")
|
||||
{
|
||||
tracing::info!(round = round + 1, "Cleanup: save dialog detected, clicking Don't Save");
|
||||
if mouse_ops::click_element(&self.kvm, "Don't Save", &scene, 100).await
|
||||
|| mouse_ops::click_element(&self.kvm, "不保存", &scene, 100).await
|
||||
|| mouse_ops::click_element(&self.kvm, "No", &scene, 100).await
|
||||
{
|
||||
sleep(Duration::from_millis(500)).await;
|
||||
continue;
|
||||
}
|
||||
}
|
||||
|
||||
// Check if we're at desktop (taskbar visible, no prominent app window)
|
||||
let has_taskbar = text_lower.contains("start")
|
||||
|| text_lower.contains("开始")
|
||||
|| text_lower.contains("search")
|
||||
|| text_lower.contains("搜索");
|
||||
|
||||
if has_taskbar {
|
||||
tracing::info!(round = round + 1, "Cleanup: desktop reached");
|
||||
return;
|
||||
}
|
||||
|
||||
// Otherwise try closing whatever window is on top
|
||||
tracing::info!(round = round + 1, "Cleanup: closing top window");
|
||||
let _ = mouse_ops::close_window(&self.kvm, &scene).await;
|
||||
sleep(Duration::from_millis(500)).await;
|
||||
}
|
||||
|
||||
tracing::info!(rounds = max_rounds, "Cleanup: max rounds reached");
|
||||
}
|
||||
}
|
||||
|
||||
/// Send Space to wake a sleeping target PC, then verify screen changed.
|
||||
|
||||
@@ -2,13 +2,17 @@
|
||||
|
||||
use agent_config::AgentConfig;
|
||||
use agent_core::KvmAgent;
|
||||
use api_server::{AppState, PlannerMetrics, RuntimeConfig};
|
||||
use clap::{Parser, Subcommand};
|
||||
use hybrid_planner::HybridPlanner;
|
||||
use kvm_client::KvmClient;
|
||||
use llm_planner::LlmPlanner;
|
||||
use memory_client::MemoryClient;
|
||||
use npu_client::NpuClient;
|
||||
use runner::{AutonomousRunner, TaskQueue};
|
||||
use std::sync::Arc;
|
||||
use template_store::TemplateStore;
|
||||
use tokio::sync::{broadcast, RwLock};
|
||||
use validator::StepValidator;
|
||||
|
||||
#[derive(Parser)]
|
||||
@@ -34,15 +38,14 @@ enum Commands {
|
||||
#[arg(long)]
|
||||
kvm_url: Option<String>,
|
||||
},
|
||||
/// Run as a daemon polling a task queue.
|
||||
/// Run as a daemon: API server + task queue polling.
|
||||
Daemon,
|
||||
/// Start the API server (port 8890).
|
||||
/// Start only the API server (port 8890), no queue polling.
|
||||
Serve,
|
||||
}
|
||||
|
||||
#[tokio::main]
|
||||
async fn main() -> anyhow::Result<()> {
|
||||
// Initialize tracing
|
||||
tracing_subscriber::fmt()
|
||||
.with_env_filter(
|
||||
tracing_subscriber::EnvFilter::try_from_default_env()
|
||||
@@ -51,8 +54,6 @@ async fn main() -> anyhow::Result<()> {
|
||||
.init();
|
||||
|
||||
let cli = Cli::parse();
|
||||
|
||||
// Load config
|
||||
let config = AgentConfig::from_yaml(std::path::Path::new(&cli.config));
|
||||
|
||||
match cli.command {
|
||||
@@ -61,21 +62,32 @@ async fn main() -> anyhow::Result<()> {
|
||||
run_single_task(&config, url, &task).await?;
|
||||
}
|
||||
Commands::Daemon => {
|
||||
tracing::info!("Daemon mode not yet implemented");
|
||||
run_daemon(config).await?;
|
||||
}
|
||||
Commands::Serve => {
|
||||
tracing::info!("API server mode not yet implemented");
|
||||
run_serve(config).await?;
|
||||
}
|
||||
}
|
||||
|
||||
Ok(())
|
||||
}
|
||||
|
||||
async fn run_single_task(config: &AgentConfig, kvm_url: &str, task: &str) -> anyhow::Result<()> {
|
||||
// Build KVM client
|
||||
/// Build the MariaDB connection URL from config + env vars.
|
||||
fn build_database_url(config: &AgentConfig) -> String {
|
||||
let pass = std::env::var("KVM_AGENT_DB_PASS").unwrap_or_else(|_| "changeme".into());
|
||||
format!(
|
||||
"mysql://{}:{}@{}/{}",
|
||||
config.db_user, pass, config.db_host, config.db_name
|
||||
)
|
||||
}
|
||||
|
||||
/// Build all subsystems needed for agent execution.
|
||||
fn build_subsystems(
|
||||
config: &AgentConfig,
|
||||
kvm_url: &str,
|
||||
) -> anyhow::Result<(KvmClient, HybridPlanner, Option<StepValidator>)> {
|
||||
let kvm = KvmClient::new(kvm_url, &config.kvm_jwt_token, config.timeout as u64)?;
|
||||
|
||||
// Build LLM planner
|
||||
let llm = LlmPlanner::new(
|
||||
&config.llm_base_url,
|
||||
&config.llm_api_key,
|
||||
@@ -86,7 +98,6 @@ async fn run_single_task(config: &AgentConfig, kvm_url: &str, task: &str) -> any
|
||||
config.adaptive_image,
|
||||
);
|
||||
|
||||
// Build optional components
|
||||
let memory = if config.memory_enabled && !config.memory_base_url.is_empty() {
|
||||
MemoryClient::new(&config.memory_base_url, config.timeout as u64).ok()
|
||||
} else {
|
||||
@@ -106,7 +117,6 @@ async fn run_single_task(config: &AgentConfig, kvm_url: &str, task: &str) -> any
|
||||
config.template_min_reliability,
|
||||
)
|
||||
});
|
||||
let validator = Some(StepValidator::default());
|
||||
|
||||
let planner = HybridPlanner::new(
|
||||
llm,
|
||||
@@ -119,7 +129,13 @@ async fn run_single_task(config: &AgentConfig, kvm_url: &str, task: &str) -> any
|
||||
config.privacy_redact_types_vec(),
|
||||
);
|
||||
|
||||
Ok((kvm, planner, Some(StepValidator::default())))
|
||||
}
|
||||
|
||||
async fn run_single_task(config: &AgentConfig, kvm_url: &str, task: &str) -> anyhow::Result<()> {
|
||||
let (kvm, planner, validator) = build_subsystems(config, kvm_url)?;
|
||||
let mut agent = KvmAgent::new(kvm, planner, Some(config.clone()), validator);
|
||||
|
||||
tracing::info!(task, "Starting single task");
|
||||
let result = agent.run_task(task).await;
|
||||
|
||||
@@ -136,6 +152,97 @@ async fn run_single_task(config: &AgentConfig, kvm_url: &str, task: &str) -> any
|
||||
"Task failed",
|
||||
);
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
|
||||
/// Daemon mode: run API server on port 8890 + poll task queue concurrently.
|
||||
async fn run_daemon(config: AgentConfig) -> anyhow::Result<()> {
|
||||
let db_url = build_database_url(&config);
|
||||
tracing::info!(db_host = config.db_host.as_str(), "Connecting to MariaDB");
|
||||
|
||||
let queue = TaskQueue::connect(&db_url).await?;
|
||||
let (events_tx, _) = broadcast::channel(256);
|
||||
|
||||
let state = Arc::new(AppState {
|
||||
queue: queue.clone(),
|
||||
config: RwLock::new(RuntimeConfig {
|
||||
llm_model: config.llm_model.clone(),
|
||||
max_steps: config.max_steps,
|
||||
llm_temperature: config.llm_temperature,
|
||||
step_delay: config.step_delay,
|
||||
}),
|
||||
metrics: RwLock::new(PlannerMetrics::default()),
|
||||
events_tx: events_tx.clone(),
|
||||
});
|
||||
|
||||
let runner = AutonomousRunner::new(queue, config.daemon_poll_interval as u64);
|
||||
|
||||
// Clone what the runner closure needs
|
||||
let config_for_runner = config.clone();
|
||||
let state_for_runner = state.clone();
|
||||
|
||||
tracing::info!("Starting daemon: API server (8890) + task queue poller");
|
||||
|
||||
tokio::select! {
|
||||
res = api_server::serve(state, 8890) => {
|
||||
tracing::error!("API server exited: {:?}", res);
|
||||
}
|
||||
res = runner.run(|task| {
|
||||
let cfg = config_for_runner.clone();
|
||||
let st = state_for_runner.clone();
|
||||
async move {
|
||||
st.emit("task_started", serde_json::json!({
|
||||
"id": task.id,
|
||||
"description": &task.description,
|
||||
}));
|
||||
|
||||
let result = match build_subsystems(&cfg, &cfg.kvm_url) {
|
||||
Ok((kvm, planner, validator)) => {
|
||||
let mut agent = KvmAgent::new(kvm, planner, Some(cfg.clone()), validator);
|
||||
agent.run_task(&task.description).await
|
||||
}
|
||||
Err(e) => {
|
||||
tracing::error!(error = %e, "Failed to build agent subsystems");
|
||||
return (false, format!("Init error: {e}"));
|
||||
}
|
||||
};
|
||||
|
||||
st.emit("task_completed", serde_json::json!({
|
||||
"id": task.id,
|
||||
"success": result.success,
|
||||
"steps": result.steps_taken,
|
||||
"reason": &result.final_reason,
|
||||
}));
|
||||
|
||||
(result.success, result.final_reason)
|
||||
}
|
||||
}) => {
|
||||
tracing::error!("Runner exited: {:?}", res);
|
||||
}
|
||||
}
|
||||
|
||||
Ok(())
|
||||
}
|
||||
|
||||
/// Serve-only mode: API server without queue polling.
|
||||
async fn run_serve(config: AgentConfig) -> anyhow::Result<()> {
|
||||
let db_url = build_database_url(&config);
|
||||
let queue = TaskQueue::connect(&db_url).await?;
|
||||
let (events_tx, _) = broadcast::channel(256);
|
||||
|
||||
let state = Arc::new(AppState {
|
||||
queue,
|
||||
config: RwLock::new(RuntimeConfig {
|
||||
llm_model: config.llm_model.clone(),
|
||||
max_steps: config.max_steps,
|
||||
llm_temperature: config.llm_temperature,
|
||||
step_delay: config.step_delay,
|
||||
}),
|
||||
metrics: RwLock::new(PlannerMetrics::default()),
|
||||
events_tx,
|
||||
});
|
||||
|
||||
tracing::info!("Starting API server on port 8890 (serve-only, no queue polling)");
|
||||
api_server::serve(state, 8890).await?;
|
||||
Ok(())
|
||||
}
|
||||
|
||||
@@ -0,0 +1,17 @@
|
||||
[package]
|
||||
name = "api-server"
|
||||
version = "0.1.0"
|
||||
edition = "2021"
|
||||
|
||||
[dependencies]
|
||||
agent-types = { workspace = true }
|
||||
agent-config = { workspace = true }
|
||||
runner = { workspace = true }
|
||||
serde = { workspace = true }
|
||||
serde_json = { workspace = true }
|
||||
tokio = { workspace = true }
|
||||
tracing = { workspace = true }
|
||||
thiserror = { workspace = true }
|
||||
axum = { workspace = true }
|
||||
tower-http = { workspace = true }
|
||||
tokio-stream = { workspace = true }
|
||||
@@ -0,0 +1,282 @@
|
||||
//! KVM Agent REST API + SSE event stream (port 8890).
|
||||
//!
|
||||
//! Endpoints:
|
||||
//! GET /api/v1/agent/status — queue status
|
||||
//! POST /api/v1/agent/tasks — enqueue new task
|
||||
//! GET /api/v1/agent/tasks — list tasks (?status=&limit=)
|
||||
//! GET /api/v1/agent/tasks/:id — get task by ID
|
||||
//! DELETE /api/v1/agent/tasks/:id — delete pending task
|
||||
//! GET /api/v1/agent/config — current config (safe subset)
|
||||
//! PATCH /api/v1/agent/config — update config fields
|
||||
//! GET /api/v1/agent/metrics — planner routing stats
|
||||
//! GET /api/v1/agent/events — SSE event stream
|
||||
|
||||
use axum::{
|
||||
extract::{Path, Query, State},
|
||||
http::StatusCode,
|
||||
response::{
|
||||
sse::{Event, KeepAlive, Sse},
|
||||
IntoResponse, Json,
|
||||
},
|
||||
routing::{get, post},
|
||||
Router,
|
||||
};
|
||||
use runner::{CreateTask, TaskQueue};
|
||||
use serde::{Deserialize, Serialize};
|
||||
use std::sync::Arc;
|
||||
use tokio::sync::{broadcast, RwLock};
|
||||
use tokio_stream::wrappers::BroadcastStream;
|
||||
use tokio_stream::StreamExt;
|
||||
use tower_http::cors::{Any, CorsLayer};
|
||||
|
||||
/// Mutable config fields that can be patched at runtime.
|
||||
#[derive(Debug, Clone, Serialize, Deserialize)]
|
||||
pub struct RuntimeConfig {
|
||||
pub llm_model: String,
|
||||
pub max_steps: u32,
|
||||
pub llm_temperature: f64,
|
||||
pub step_delay: f64,
|
||||
}
|
||||
|
||||
/// Planner routing metrics (matches Python PlannerMetrics).
|
||||
#[derive(Debug, Clone, Default, Serialize, Deserialize)]
|
||||
pub struct PlannerMetrics {
|
||||
pub template_hits: u64,
|
||||
pub fingerprint_hits: u64,
|
||||
pub local_llm_calls: u64,
|
||||
pub remote_llm_calls: u64,
|
||||
pub privacy_redactions: u64,
|
||||
pub privacy_text_only_mode: u64,
|
||||
pub privacy_image_mode: u64,
|
||||
}
|
||||
|
||||
/// SSE event payload.
|
||||
#[derive(Debug, Clone, Serialize)]
|
||||
pub struct AgentEvent {
|
||||
pub event_type: String,
|
||||
pub data: serde_json::Value,
|
||||
}
|
||||
|
||||
/// Shared application state.
|
||||
pub struct AppState {
|
||||
pub queue: TaskQueue,
|
||||
pub config: RwLock<RuntimeConfig>,
|
||||
pub metrics: RwLock<PlannerMetrics>,
|
||||
pub events_tx: broadcast::Sender<AgentEvent>,
|
||||
}
|
||||
|
||||
impl AppState {
|
||||
/// Broadcast an event to all SSE listeners.
|
||||
pub fn emit(&self, event_type: &str, data: serde_json::Value) {
|
||||
let _ = self.events_tx.send(AgentEvent {
|
||||
event_type: event_type.into(),
|
||||
data,
|
||||
});
|
||||
}
|
||||
}
|
||||
|
||||
/// Build the axum Router with all agent API routes.
|
||||
pub fn build_router(state: Arc<AppState>) -> Router {
|
||||
let cors = CorsLayer::new()
|
||||
.allow_origin(Any)
|
||||
.allow_methods(Any)
|
||||
.allow_headers(Any);
|
||||
|
||||
let api = Router::new()
|
||||
.route("/status", get(get_status))
|
||||
.route("/tasks", post(create_task).get(list_tasks))
|
||||
.route("/tasks/{id}", get(get_task).delete(delete_task))
|
||||
.route("/config", get(get_config).patch(patch_config))
|
||||
.route("/metrics", get(get_metrics))
|
||||
.route("/events", get(sse_events));
|
||||
|
||||
Router::new()
|
||||
.nest("/api/v1/agent", api)
|
||||
.layer(cors)
|
||||
.with_state(state)
|
||||
}
|
||||
|
||||
/// Start the API server on the given port.
|
||||
pub async fn serve(state: Arc<AppState>, port: u16) -> Result<(), std::io::Error> {
|
||||
let app = build_router(state);
|
||||
let listener = tokio::net::TcpListener::bind(format!("0.0.0.0:{port}")).await?;
|
||||
tracing::info!(port, "API server listening");
|
||||
axum::serve(listener, app).await
|
||||
}
|
||||
|
||||
// ── Handlers ─────────────────────────────────────────────────
|
||||
|
||||
async fn get_status(State(state): State<Arc<AppState>>) -> impl IntoResponse {
|
||||
match state.queue.status().await {
|
||||
Ok(s) => Json(s).into_response(),
|
||||
Err(e) => (StatusCode::INTERNAL_SERVER_ERROR, e.to_string()).into_response(),
|
||||
}
|
||||
}
|
||||
|
||||
async fn create_task(
|
||||
State(state): State<Arc<AppState>>,
|
||||
Json(req): Json<CreateTask>,
|
||||
) -> impl IntoResponse {
|
||||
match state.queue.create(&req).await {
|
||||
Ok(id) => {
|
||||
state.emit(
|
||||
"task_created",
|
||||
serde_json::json!({"id": id, "description": req.description}),
|
||||
);
|
||||
(
|
||||
StatusCode::CREATED,
|
||||
Json(serde_json::json!({"id": id, "status": "pending"})),
|
||||
)
|
||||
.into_response()
|
||||
}
|
||||
Err(e) => (StatusCode::INTERNAL_SERVER_ERROR, e.to_string()).into_response(),
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Deserialize)]
|
||||
struct ListQuery {
|
||||
status: Option<String>,
|
||||
limit: Option<u32>,
|
||||
}
|
||||
|
||||
async fn list_tasks(
|
||||
State(state): State<Arc<AppState>>,
|
||||
Query(q): Query<ListQuery>,
|
||||
) -> impl IntoResponse {
|
||||
let limit = q.limit.unwrap_or(50).min(200);
|
||||
match state.queue.list(q.status.as_deref(), limit).await {
|
||||
Ok(tasks) => Json(tasks).into_response(),
|
||||
Err(e) => (StatusCode::INTERNAL_SERVER_ERROR, e.to_string()).into_response(),
|
||||
}
|
||||
}
|
||||
|
||||
async fn get_task(
|
||||
State(state): State<Arc<AppState>>,
|
||||
Path(id): Path<i64>,
|
||||
) -> impl IntoResponse {
|
||||
match state.queue.get(id).await {
|
||||
Ok(t) => Json(t).into_response(),
|
||||
Err(runner::RunnerError::NotFound(_)) => StatusCode::NOT_FOUND.into_response(),
|
||||
Err(e) => (StatusCode::INTERNAL_SERVER_ERROR, e.to_string()).into_response(),
|
||||
}
|
||||
}
|
||||
|
||||
async fn delete_task(
|
||||
State(state): State<Arc<AppState>>,
|
||||
Path(id): Path<i64>,
|
||||
) -> impl IntoResponse {
|
||||
match state.queue.delete(id).await {
|
||||
Ok(()) => StatusCode::NO_CONTENT.into_response(),
|
||||
Err(runner::RunnerError::NotFound(_)) => StatusCode::NOT_FOUND.into_response(),
|
||||
Err(e) => (StatusCode::INTERNAL_SERVER_ERROR, e.to_string()).into_response(),
|
||||
}
|
||||
}
|
||||
|
||||
async fn get_config(State(state): State<Arc<AppState>>) -> impl IntoResponse {
|
||||
let cfg = state.config.read().await;
|
||||
Json(cfg.clone())
|
||||
}
|
||||
|
||||
#[derive(Deserialize)]
|
||||
struct PatchConfig {
|
||||
llm_model: Option<String>,
|
||||
max_steps: Option<u32>,
|
||||
llm_temperature: Option<f64>,
|
||||
step_delay: Option<f64>,
|
||||
}
|
||||
|
||||
async fn patch_config(
|
||||
State(state): State<Arc<AppState>>,
|
||||
Json(patch): Json<PatchConfig>,
|
||||
) -> impl IntoResponse {
|
||||
let mut cfg = state.config.write().await;
|
||||
if let Some(m) = patch.llm_model {
|
||||
cfg.llm_model = m;
|
||||
}
|
||||
if let Some(s) = patch.max_steps {
|
||||
cfg.max_steps = s;
|
||||
}
|
||||
if let Some(t) = patch.llm_temperature {
|
||||
cfg.llm_temperature = t;
|
||||
}
|
||||
if let Some(d) = patch.step_delay {
|
||||
cfg.step_delay = d;
|
||||
}
|
||||
state.emit("config_changed", serde_json::to_value(&*cfg).unwrap());
|
||||
Json(cfg.clone())
|
||||
}
|
||||
|
||||
async fn get_metrics(State(state): State<Arc<AppState>>) -> impl IntoResponse {
|
||||
let m = state.metrics.read().await;
|
||||
Json(m.clone())
|
||||
}
|
||||
|
||||
async fn sse_events(
|
||||
State(state): State<Arc<AppState>>,
|
||||
) -> Sse<impl tokio_stream::Stream<Item = Result<Event, std::convert::Infallible>>> {
|
||||
let rx = state.events_tx.subscribe();
|
||||
let stream = BroadcastStream::new(rx).filter_map(|msg| match msg {
|
||||
Ok(evt) => {
|
||||
let data = serde_json::to_string(&evt.data).unwrap_or_default();
|
||||
Some(Ok(Event::default().event(evt.event_type).data(data)))
|
||||
}
|
||||
Err(_) => None, // lagged receiver — skip
|
||||
});
|
||||
|
||||
Sse::new(stream).keep_alive(KeepAlive::default())
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
|
||||
#[test]
|
||||
fn test_runtime_config_serde() {
|
||||
let cfg = RuntimeConfig {
|
||||
llm_model: "gpt-4o".into(),
|
||||
max_steps: 30,
|
||||
llm_temperature: 0.3,
|
||||
step_delay: 1.0,
|
||||
};
|
||||
let json = serde_json::to_string(&cfg).unwrap();
|
||||
let parsed: RuntimeConfig = serde_json::from_str(&json).unwrap();
|
||||
assert_eq!(parsed.llm_model, "gpt-4o");
|
||||
assert_eq!(parsed.max_steps, 30);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_planner_metrics_default() {
|
||||
let m = PlannerMetrics::default();
|
||||
assert_eq!(m.template_hits, 0);
|
||||
assert_eq!(m.remote_llm_calls, 0);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_agent_event_serde() {
|
||||
let evt = AgentEvent {
|
||||
event_type: "task_created".into(),
|
||||
data: serde_json::json!({"id": 1}),
|
||||
};
|
||||
let json = serde_json::to_string(&evt).unwrap();
|
||||
assert!(json.contains("task_created"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_broadcast_no_panic_without_receivers() {
|
||||
let (tx, _rx) = broadcast::channel::<AgentEvent>(16);
|
||||
// Sending to a channel with no active receivers should not panic
|
||||
let _ = tx.send(AgentEvent {
|
||||
event_type: "test".into(),
|
||||
data: serde_json::json!({}),
|
||||
});
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_patch_config_partial() {
|
||||
let patch: PatchConfig = serde_json::from_str(r#"{"max_steps": 50}"#).unwrap();
|
||||
assert_eq!(patch.max_steps, Some(50));
|
||||
assert!(patch.llm_model.is_none());
|
||||
assert!(patch.llm_temperature.is_none());
|
||||
assert!(patch.step_delay.is_none());
|
||||
}
|
||||
}
|
||||
@@ -28,12 +28,21 @@ pub struct KvmClient {
|
||||
pub struct OcrSnapshot {
|
||||
#[serde(default)]
|
||||
pub text: String,
|
||||
#[serde(default)]
|
||||
#[serde(default, deserialize_with = "deserialize_null_as_empty")]
|
||||
pub regions: Vec<OcrRegion>,
|
||||
#[serde(default)]
|
||||
pub processing_ms: u64,
|
||||
}
|
||||
|
||||
/// Deserialize `null` or missing field as empty Vec.
|
||||
fn deserialize_null_as_empty<'de, D, T>(deserializer: D) -> std::result::Result<Vec<T>, D::Error>
|
||||
where
|
||||
D: serde::Deserializer<'de>,
|
||||
T: serde::Deserialize<'de>,
|
||||
{
|
||||
Option::<Vec<T>>::deserialize(deserializer).map(|opt| opt.unwrap_or_default())
|
||||
}
|
||||
|
||||
#[derive(Debug, Deserialize)]
|
||||
pub struct OcrRegion {
|
||||
pub text: String,
|
||||
@@ -365,6 +374,22 @@ fn shortcut_to_kvm_name(shortcut: &str) -> String {
|
||||
mod tests {
|
||||
use super::*;
|
||||
|
||||
#[test]
|
||||
fn ocr_snapshot_null_regions() {
|
||||
let json = r#"{"text":"hello","regions":null,"processing_ms":120}"#;
|
||||
let snap: OcrSnapshot = serde_json::from_str(json).unwrap();
|
||||
assert_eq!(snap.text, "hello");
|
||||
assert!(snap.regions.is_empty());
|
||||
assert_eq!(snap.processing_ms, 120);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn ocr_snapshot_missing_regions() {
|
||||
let json = r#"{"text":"hello","processing_ms":50}"#;
|
||||
let snap: OcrSnapshot = serde_json::from_str(json).unwrap();
|
||||
assert!(snap.regions.is_empty());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn shortcut_mapping() {
|
||||
assert_eq!(shortcut_to_kvm_name("Ctrl+C"), "ctrl_c");
|
||||
|
||||
@@ -179,7 +179,11 @@ impl LlmPlanner {
|
||||
adaptive_image: bool,
|
||||
) -> Self {
|
||||
Self {
|
||||
client: reqwest::Client::new(),
|
||||
client: reqwest::Client::builder()
|
||||
.timeout(std::time::Duration::from_secs(60))
|
||||
.connect_timeout(std::time::Duration::from_secs(10))
|
||||
.build()
|
||||
.unwrap_or_else(|_| reqwest::Client::new()),
|
||||
base_url: base_url.trim_end_matches('/').to_string(),
|
||||
api_key: api_key.to_string(),
|
||||
model: model.to_string(),
|
||||
@@ -209,7 +213,19 @@ impl LlmPlanner {
|
||||
}
|
||||
}
|
||||
|
||||
/// Returns true if the planner has a valid base URL configured.
|
||||
pub fn is_configured(&self) -> bool {
|
||||
!self.base_url.is_empty()
|
||||
}
|
||||
|
||||
async fn try_plan_action(&self, req: &PlanRequest<'_>) -> Result<Action, PlannerError> {
|
||||
if self.base_url.is_empty() {
|
||||
return Ok(Action::Wait {
|
||||
delay: 2.0,
|
||||
reason: "LLM not configured: base_url is empty".into(),
|
||||
});
|
||||
}
|
||||
|
||||
let mut user_parts: Vec<serde_json::Value> = Vec::new();
|
||||
|
||||
// Task + step
|
||||
@@ -293,9 +309,16 @@ impl LlmPlanner {
|
||||
temperature: self.temperature,
|
||||
};
|
||||
|
||||
// base_url may or may not include "/v1" — normalize
|
||||
let url = if self.base_url.ends_with("/v1") || self.base_url.ends_with("/v1/") {
|
||||
format!("{}/chat/completions", self.base_url.trim_end_matches('/'))
|
||||
} else {
|
||||
format!("{}/v1/chat/completions", self.base_url)
|
||||
};
|
||||
|
||||
let resp = self
|
||||
.client
|
||||
.post(format!("{}/v1/chat/completions", self.base_url))
|
||||
.post(&url)
|
||||
.bearer_auth(&self.api_key)
|
||||
.json(&chat_req)
|
||||
.send()
|
||||
@@ -360,7 +383,32 @@ pub fn parse_action(content: &str) -> Action {
|
||||
let raw: RawAction = match serde_json::from_str(json_text) {
|
||||
Ok(r) => r,
|
||||
Err(_) => {
|
||||
tracing::warn!(content = &content[..content.len().min(200)], "Failed to parse LLM response as JSON");
|
||||
// If the response looks like a natural language description
|
||||
// (common for "describe" tasks), wrap it as a Done action.
|
||||
let trimmed = content.trim();
|
||||
if trimmed.len() > 50 && !trimmed.starts_with('{') {
|
||||
tracing::info!("LLM returned natural language — wrapping as Done");
|
||||
let msg = if trimmed.len() > 500 {
|
||||
// Find a valid UTF-8 char boundary near 500 bytes
|
||||
let end = trimmed
|
||||
.char_indices()
|
||||
.take_while(|(i, _)| *i <= 500)
|
||||
.last()
|
||||
.map(|(i, c)| i + c.len_utf8())
|
||||
.unwrap_or(500.min(trimmed.len()));
|
||||
format!("{}...", &trimmed[..end])
|
||||
} else {
|
||||
trimmed.to_string()
|
||||
};
|
||||
return Action::Done { message: msg };
|
||||
}
|
||||
let preview_end = content
|
||||
.char_indices()
|
||||
.take_while(|(i, _)| *i <= 200)
|
||||
.last()
|
||||
.map(|(i, c)| i + c.len_utf8())
|
||||
.unwrap_or(content.len().min(200));
|
||||
tracing::warn!(content = &content[..preview_end], "Failed to parse LLM response as JSON");
|
||||
return Action::Wait {
|
||||
delay: 1.0,
|
||||
reason: "Failed to parse LLM response".into(),
|
||||
|
||||
@@ -0,0 +1,14 @@
|
||||
[package]
|
||||
name = "runner"
|
||||
version = "0.1.0"
|
||||
edition = "2021"
|
||||
|
||||
[dependencies]
|
||||
agent-types = { workspace = true }
|
||||
agent-config = { workspace = true }
|
||||
serde = { workspace = true }
|
||||
serde_json = { workspace = true }
|
||||
tokio = { workspace = true }
|
||||
tracing = { workspace = true }
|
||||
thiserror = { workspace = true }
|
||||
sqlx = { workspace = true }
|
||||
@@ -0,0 +1,388 @@
|
||||
//! Task queue backed by MariaDB — atomic dequeue with `FOR UPDATE`.
|
||||
|
||||
use serde::{Deserialize, Serialize};
|
||||
use sqlx::mysql::{MySqlPool, MySqlPoolOptions};
|
||||
use std::time::{SystemTime, UNIX_EPOCH};
|
||||
use tokio::time::{sleep, Duration};
|
||||
use tracing;
|
||||
|
||||
#[derive(thiserror::Error, Debug)]
|
||||
pub enum RunnerError {
|
||||
#[error("database error: {0}")]
|
||||
Db(#[from] sqlx::Error),
|
||||
#[error("task not found: {0}")]
|
||||
NotFound(i64),
|
||||
}
|
||||
|
||||
/// Task row matching the `agent_tasks` table schema.
|
||||
#[derive(Debug, Clone, Serialize, Deserialize, sqlx::FromRow)]
|
||||
pub struct AgentTask {
|
||||
pub id: i64,
|
||||
pub task_type: String,
|
||||
pub description: String,
|
||||
pub variables: serde_json::Value,
|
||||
pub status: String,
|
||||
pub created_at: f64,
|
||||
pub started_at: Option<f64>,
|
||||
pub finished_at: Option<f64>,
|
||||
pub result: Option<String>,
|
||||
pub success: Option<i8>,
|
||||
}
|
||||
|
||||
/// Lightweight view returned by list queries.
|
||||
#[derive(Debug, Clone, Serialize, Deserialize)]
|
||||
pub struct TaskSummary {
|
||||
pub id: i64,
|
||||
pub task_type: String,
|
||||
pub description: String,
|
||||
pub status: String,
|
||||
pub created_at: f64,
|
||||
pub success: Option<i8>,
|
||||
}
|
||||
|
||||
impl From<AgentTask> for TaskSummary {
|
||||
fn from(t: AgentTask) -> Self {
|
||||
Self {
|
||||
id: t.id,
|
||||
task_type: t.task_type,
|
||||
description: t.description,
|
||||
status: t.status,
|
||||
created_at: t.created_at,
|
||||
success: t.success,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/// New task request.
|
||||
#[derive(Debug, Clone, Serialize, Deserialize)]
|
||||
pub struct CreateTask {
|
||||
#[serde(default = "default_task_type")]
|
||||
pub task_type: String,
|
||||
pub description: String,
|
||||
#[serde(default)]
|
||||
pub variables: serde_json::Value,
|
||||
}
|
||||
|
||||
fn default_task_type() -> String {
|
||||
"task".into()
|
||||
}
|
||||
|
||||
fn now_epoch() -> f64 {
|
||||
SystemTime::now()
|
||||
.duration_since(UNIX_EPOCH)
|
||||
.unwrap_or_default()
|
||||
.as_secs_f64()
|
||||
}
|
||||
|
||||
/// MariaDB-backed task queue.
|
||||
#[derive(Clone)]
|
||||
pub struct TaskQueue {
|
||||
pool: MySqlPool,
|
||||
}
|
||||
|
||||
impl TaskQueue {
|
||||
/// Connect to MariaDB and ensure the `agent_tasks` table exists.
|
||||
pub async fn connect(database_url: &str) -> Result<Self, RunnerError> {
|
||||
let pool = MySqlPoolOptions::new()
|
||||
.max_connections(5)
|
||||
.connect(database_url)
|
||||
.await?;
|
||||
|
||||
// Auto-create table if missing
|
||||
sqlx::query(
|
||||
r#"CREATE TABLE IF NOT EXISTS agent_tasks (
|
||||
id BIGINT AUTO_INCREMENT PRIMARY KEY,
|
||||
task_type VARCHAR(50) NOT NULL DEFAULT 'task',
|
||||
description TEXT NOT NULL,
|
||||
variables JSON NOT NULL,
|
||||
status VARCHAR(20) NOT NULL DEFAULT 'pending',
|
||||
created_at DOUBLE NOT NULL,
|
||||
started_at DOUBLE,
|
||||
finished_at DOUBLE,
|
||||
result TEXT DEFAULT '',
|
||||
success TINYINT(1) DEFAULT 0,
|
||||
INDEX idx_status (status),
|
||||
INDEX idx_created (created_at)
|
||||
) ENGINE=InnoDB DEFAULT CHARSET=utf8mb4"#,
|
||||
)
|
||||
.execute(&pool)
|
||||
.await?;
|
||||
|
||||
Ok(Self { pool })
|
||||
}
|
||||
|
||||
/// Insert a new pending task, return its ID.
|
||||
pub async fn create(&self, req: &CreateTask) -> Result<i64, RunnerError> {
|
||||
let vars = serde_json::to_string(&req.variables).unwrap_or_else(|_| "{}".into());
|
||||
let now = now_epoch();
|
||||
|
||||
let result = sqlx::query(
|
||||
"INSERT INTO agent_tasks (task_type, description, variables, status, created_at) VALUES (?, ?, ?, 'pending', ?)"
|
||||
)
|
||||
.bind(&req.task_type)
|
||||
.bind(&req.description)
|
||||
.bind(&vars)
|
||||
.bind(now)
|
||||
.execute(&self.pool)
|
||||
.await?;
|
||||
|
||||
Ok(result.last_insert_id() as i64)
|
||||
}
|
||||
|
||||
/// Atomically claim the oldest pending task.
|
||||
pub async fn dequeue(&self) -> Result<Option<AgentTask>, RunnerError> {
|
||||
let mut tx = self.pool.begin().await?;
|
||||
|
||||
let row: Option<AgentTask> = sqlx::query_as(
|
||||
"SELECT * FROM agent_tasks WHERE status = 'pending' ORDER BY created_at ASC LIMIT 1 FOR UPDATE"
|
||||
)
|
||||
.fetch_optional(&mut *tx)
|
||||
.await?;
|
||||
|
||||
if let Some(ref task) = row {
|
||||
let now = now_epoch();
|
||||
sqlx::query("UPDATE agent_tasks SET status = 'running', started_at = ? WHERE id = ?")
|
||||
.bind(now)
|
||||
.bind(task.id)
|
||||
.execute(&mut *tx)
|
||||
.await?;
|
||||
}
|
||||
|
||||
tx.commit().await?;
|
||||
|
||||
// Return a copy with updated status
|
||||
Ok(row.map(|mut t| {
|
||||
t.status = "running".into();
|
||||
t.started_at = Some(now_epoch());
|
||||
t
|
||||
}))
|
||||
}
|
||||
|
||||
/// Mark a task as completed.
|
||||
pub async fn complete(
|
||||
&self,
|
||||
id: i64,
|
||||
success: bool,
|
||||
result_text: &str,
|
||||
) -> Result<(), RunnerError> {
|
||||
let now = now_epoch();
|
||||
let rows = sqlx::query(
|
||||
"UPDATE agent_tasks SET status = 'done', finished_at = ?, result = ?, success = ? WHERE id = ?",
|
||||
)
|
||||
.bind(now)
|
||||
.bind(result_text)
|
||||
.bind(success as i8)
|
||||
.bind(id)
|
||||
.execute(&self.pool)
|
||||
.await?;
|
||||
|
||||
if rows.rows_affected() == 0 {
|
||||
return Err(RunnerError::NotFound(id));
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
|
||||
/// Mark a running task as failed.
|
||||
pub async fn fail(&self, id: i64, error: &str) -> Result<(), RunnerError> {
|
||||
self.complete(id, false, error).await
|
||||
}
|
||||
|
||||
/// Get a single task by ID.
|
||||
pub async fn get(&self, id: i64) -> Result<AgentTask, RunnerError> {
|
||||
sqlx::query_as("SELECT * FROM agent_tasks WHERE id = ?")
|
||||
.bind(id)
|
||||
.fetch_optional(&self.pool)
|
||||
.await?
|
||||
.ok_or(RunnerError::NotFound(id))
|
||||
}
|
||||
|
||||
/// List tasks with optional status filter and limit.
|
||||
pub async fn list(
|
||||
&self,
|
||||
status: Option<&str>,
|
||||
limit: u32,
|
||||
) -> Result<Vec<TaskSummary>, RunnerError> {
|
||||
let tasks: Vec<AgentTask> = if let Some(s) = status {
|
||||
sqlx::query_as(
|
||||
"SELECT * FROM agent_tasks WHERE status = ? ORDER BY created_at DESC LIMIT ?",
|
||||
)
|
||||
.bind(s)
|
||||
.bind(limit)
|
||||
.fetch_all(&self.pool)
|
||||
.await?
|
||||
} else {
|
||||
sqlx::query_as("SELECT * FROM agent_tasks ORDER BY created_at DESC LIMIT ?")
|
||||
.bind(limit)
|
||||
.fetch_all(&self.pool)
|
||||
.await?
|
||||
};
|
||||
|
||||
Ok(tasks.into_iter().map(TaskSummary::from).collect())
|
||||
}
|
||||
|
||||
/// Delete a task by ID. Only pending tasks can be deleted.
|
||||
pub async fn delete(&self, id: i64) -> Result<(), RunnerError> {
|
||||
let rows =
|
||||
sqlx::query("DELETE FROM agent_tasks WHERE id = ? AND status = 'pending'")
|
||||
.bind(id)
|
||||
.execute(&self.pool)
|
||||
.await?;
|
||||
|
||||
if rows.rows_affected() == 0 {
|
||||
return Err(RunnerError::NotFound(id));
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
|
||||
/// Get current queue status: running task (if any), pending count.
|
||||
pub async fn status(&self) -> Result<QueueStatus, RunnerError> {
|
||||
let running: Option<AgentTask> = sqlx::query_as(
|
||||
"SELECT * FROM agent_tasks WHERE status = 'running' ORDER BY started_at DESC LIMIT 1",
|
||||
)
|
||||
.fetch_optional(&self.pool)
|
||||
.await?;
|
||||
|
||||
let pending_count: (i64,) =
|
||||
sqlx::query_as("SELECT COUNT(*) FROM agent_tasks WHERE status = 'pending'")
|
||||
.fetch_one(&self.pool)
|
||||
.await?;
|
||||
|
||||
let status = if running.is_some() { "running" } else { "idle" };
|
||||
Ok(QueueStatus {
|
||||
running_task: running.map(|t| t.description),
|
||||
pending_count: pending_count.0 as u32,
|
||||
status: status.into(),
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, Serialize, Deserialize)]
|
||||
pub struct QueueStatus {
|
||||
pub status: String,
|
||||
pub running_task: Option<String>,
|
||||
pub pending_count: u32,
|
||||
}
|
||||
|
||||
/// Autonomous daemon that polls the task queue and invokes a callback for each task.
|
||||
pub struct AutonomousRunner {
|
||||
queue: TaskQueue,
|
||||
poll_interval: Duration,
|
||||
}
|
||||
|
||||
impl AutonomousRunner {
|
||||
pub fn new(queue: TaskQueue, poll_interval_secs: u64) -> Self {
|
||||
Self {
|
||||
queue,
|
||||
poll_interval: Duration::from_secs(poll_interval_secs),
|
||||
}
|
||||
}
|
||||
|
||||
/// Run forever, polling for tasks and calling `handler` for each one.
|
||||
/// The handler receives the task description and returns (success, result_message).
|
||||
pub async fn run<F, Fut>(&self, handler: F) -> Result<(), RunnerError>
|
||||
where
|
||||
F: Fn(AgentTask) -> Fut,
|
||||
Fut: std::future::Future<Output = (bool, String)>,
|
||||
{
|
||||
tracing::info!(
|
||||
interval_s = self.poll_interval.as_secs(),
|
||||
"Autonomous runner started",
|
||||
);
|
||||
|
||||
loop {
|
||||
match self.queue.dequeue().await {
|
||||
Ok(Some(task)) => {
|
||||
tracing::info!(
|
||||
id = task.id,
|
||||
desc = task.description.as_str(),
|
||||
"Dequeued task",
|
||||
);
|
||||
|
||||
let (success, result_text) = handler(task.clone()).await;
|
||||
|
||||
if let Err(e) = self.queue.complete(task.id, success, &result_text).await {
|
||||
tracing::error!(id = task.id, error = %e, "Failed to complete task");
|
||||
} else {
|
||||
tracing::info!(
|
||||
id = task.id,
|
||||
success,
|
||||
"Task completed",
|
||||
);
|
||||
}
|
||||
}
|
||||
Ok(None) => {
|
||||
sleep(self.poll_interval).await;
|
||||
}
|
||||
Err(e) => {
|
||||
tracing::error!(error = %e, "Dequeue failed, retrying after interval");
|
||||
sleep(self.poll_interval).await;
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
pub fn queue(&self) -> &TaskQueue {
|
||||
&self.queue
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
|
||||
#[test]
|
||||
fn test_create_task_defaults() {
|
||||
let json = r#"{"description": "open notepad"}"#;
|
||||
let req: CreateTask = serde_json::from_str(json).unwrap();
|
||||
assert_eq!(req.task_type, "task");
|
||||
assert_eq!(req.description, "open notepad");
|
||||
assert_eq!(req.variables, serde_json::json!(null));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_create_task_with_type() {
|
||||
let json = r#"{"task_type": "workflow", "description": "run backup", "variables": {"target": "/data"}}"#;
|
||||
let req: CreateTask = serde_json::from_str(json).unwrap();
|
||||
assert_eq!(req.task_type, "workflow");
|
||||
assert_eq!(req.variables["target"], "/data");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_task_summary_from_agent_task() {
|
||||
let task = AgentTask {
|
||||
id: 42,
|
||||
task_type: "task".into(),
|
||||
description: "test".into(),
|
||||
variables: serde_json::json!({}),
|
||||
status: "done".into(),
|
||||
created_at: 1000.0,
|
||||
started_at: Some(1001.0),
|
||||
finished_at: Some(1002.0),
|
||||
result: Some("ok".into()),
|
||||
success: Some(1),
|
||||
};
|
||||
let summary = TaskSummary::from(task);
|
||||
assert_eq!(summary.id, 42);
|
||||
assert_eq!(summary.status, "done");
|
||||
assert_eq!(summary.success, Some(1));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_now_epoch_reasonable() {
|
||||
let t = now_epoch();
|
||||
// Should be after 2025-01-01
|
||||
assert!(t > 1735689600.0);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_queue_status_serde() {
|
||||
let status = QueueStatus {
|
||||
status: "idle".into(),
|
||||
running_task: None,
|
||||
pending_count: 3,
|
||||
};
|
||||
let json = serde_json::to_string(&status).unwrap();
|
||||
assert!(json.contains("\"pending_count\":3"));
|
||||
assert!(json.contains("\"running_task\":null"));
|
||||
}
|
||||
}
|
||||
@@ -109,7 +109,9 @@ impl ScreenStateDetector {
|
||||
if self.hashes.len() > self.max_history {
|
||||
self.hashes.remove(0);
|
||||
}
|
||||
if self.hashes.len() >= 3 {
|
||||
if self.hashes.len() >= 3 && !has_regions {
|
||||
// Only detect sleep when OCR finds nothing — a static desktop
|
||||
// with visible text (has_regions=true) is NOT sleeping.
|
||||
let last3 = &self.hashes[self.hashes.len() - 3..];
|
||||
if last3[0] == last3[1] && last3[1] == last3[2] {
|
||||
return StateDetection {
|
||||
@@ -228,7 +230,7 @@ mod tests {
|
||||
fn sleep_detection() {
|
||||
let mut detector = ScreenStateDetector::new(5);
|
||||
let screenshot = b"identical_screenshot_data";
|
||||
// Need 3 identical screenshots
|
||||
// Need 3 identical screenshots with no OCR regions (has_regions=false)
|
||||
let r1 = detector.detect(screenshot, "", false);
|
||||
assert_ne!(r1.state, PCState::Sleep);
|
||||
let r2 = detector.detect(screenshot, "", false);
|
||||
@@ -238,6 +240,17 @@ mod tests {
|
||||
assert!(r3.confidence > 0.9);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn static_desktop_not_sleep() {
|
||||
let mut detector = ScreenStateDetector::new(5);
|
||||
let screenshot = b"identical_desktop_data";
|
||||
// 3 identical screenshots but has_regions=true → NOT sleep
|
||||
detector.detect(screenshot, "desktop icons", true);
|
||||
detector.detect(screenshot, "desktop icons", true);
|
||||
let r = detector.detect(screenshot, "desktop icons", true);
|
||||
assert_ne!(r.state, PCState::Sleep, "static desktop with OCR should not be Sleep");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn bios_detection() {
|
||||
let mut detector = ScreenStateDetector::new(5);
|
||||
|
||||
+18
-11
@@ -17,6 +17,7 @@ from typing import List, Optional
|
||||
from .actions import Action
|
||||
from .kvm_client import KVMClient
|
||||
from .llm_planner import LLMPlanner
|
||||
from .prompt_modules import build_context_header
|
||||
from .screen_state import ScreenStateDetector, PCState
|
||||
from .validator import ValidationResult
|
||||
from .workflow_hooks import EventBus, StepEvent
|
||||
@@ -235,22 +236,26 @@ class KVMAgent:
|
||||
|
||||
# Build scene_text from single scene
|
||||
scene_text = ""
|
||||
os_type_str = ""
|
||||
pc_state_str = ""
|
||||
if scene:
|
||||
scene_text = scene.to_text_summary()
|
||||
|
||||
# Inject system state context for non-APPLICATION states
|
||||
if (detection and
|
||||
detection.state.name not in ("APPLICATION", "DESKTOP")):
|
||||
scene_text = (
|
||||
f"[System State: {detection.state.name}]\n" + scene_text
|
||||
)
|
||||
# Detect OS type from OCR text
|
||||
os_type_str = self.state_detector.detect_os(
|
||||
scene.raw_ocr_text).value
|
||||
|
||||
# Inject inferred window title from topmost OCR elements
|
||||
# Build compact context header
|
||||
window_hint = _infer_window_title(scene)
|
||||
if window_hint:
|
||||
scene_text = (
|
||||
f"[Active Window: {window_hint}]\n" + scene_text
|
||||
)
|
||||
if detection:
|
||||
pc_state_str = detection.state.value
|
||||
header = build_context_header(
|
||||
os_type=os_type_str,
|
||||
state=pc_state_str,
|
||||
window_title=window_hint,
|
||||
)
|
||||
if header:
|
||||
scene_text = header + "\n" + scene_text
|
||||
|
||||
# ── Decide ───────────────────────────────────
|
||||
_t0 = time.monotonic()
|
||||
@@ -261,6 +266,8 @@ class KVMAgent:
|
||||
action_history,
|
||||
max_steps=self.max_steps,
|
||||
scene_text=scene_text,
|
||||
pc_state=pc_state_str,
|
||||
os_type=os_type_str,
|
||||
)
|
||||
if self._hooks:
|
||||
_plan_ms = (time.monotonic() - _t0) * 1000
|
||||
|
||||
@@ -135,6 +135,8 @@ class HybridPlanner:
|
||||
step: int,
|
||||
history: list[str],
|
||||
max_steps: int = 30,
|
||||
pc_state: str = "",
|
||||
os_type: str = "",
|
||||
) -> PlanResult:
|
||||
"""Decide the next action: template replay or LLM call.
|
||||
|
||||
@@ -144,6 +146,8 @@ class HybridPlanner:
|
||||
step: Current step number (0-indexed).
|
||||
history: List of action description strings.
|
||||
max_steps: Maximum steps for this task.
|
||||
pc_state: Current PC state (e.g. "desktop", "app_window").
|
||||
os_type: Detected OS type (e.g. "windows", "gnome").
|
||||
"""
|
||||
# On step 0, try to find a matching template
|
||||
if step == 0:
|
||||
@@ -213,6 +217,7 @@ class HybridPlanner:
|
||||
_t0 = time.monotonic()
|
||||
result = await self._llm_plan_with_memory(
|
||||
screenshot, task, step, history, max_steps,
|
||||
pc_state=pc_state, os_type=os_type,
|
||||
)
|
||||
if self._hooks:
|
||||
_model = getattr(self._llm, 'model', 'remote_llm')
|
||||
@@ -278,6 +283,8 @@ class HybridPlanner:
|
||||
step: int,
|
||||
history: list[str],
|
||||
max_steps: int,
|
||||
pc_state: str = "",
|
||||
os_type: str = "",
|
||||
) -> PlanResult:
|
||||
"""Plan via remote LLM with PII redaction + memory context.
|
||||
|
||||
@@ -332,6 +339,8 @@ class HybridPlanner:
|
||||
max_steps=max_steps,
|
||||
scene_text=scene_text,
|
||||
context_hint=context_hint,
|
||||
pc_state=pc_state,
|
||||
os_type=os_type,
|
||||
)
|
||||
|
||||
# Store turn in memory
|
||||
@@ -437,15 +446,17 @@ class HybridPlanner:
|
||||
scene = await self._perception.perceive(screenshot)
|
||||
scene_text = scene.to_text_summary()
|
||||
|
||||
# Compact prompt for small model (1.5B context window is limited)
|
||||
# Structured few-shot prompt for small model (1.5B, ~4K context)
|
||||
prompt = (
|
||||
f"Screen Elements:\n{scene_text[:800]}\n\n"
|
||||
f"Task: {task}\n"
|
||||
f"Step {step}/{max_steps}. History: {'; '.join(history[-3:]) if history else 'none'}\n"
|
||||
f"Screen: {scene_text[:1000]}\n"
|
||||
"Respond with ONE JSON action: "
|
||||
'{"type":"click","x":0.5,"y":0.5,"reason":"..."} or '
|
||||
'{"type":"type","text":"...","reason":"..."} or '
|
||||
'{"type":"done","reason":"..."}'
|
||||
f"Step {step}/{max_steps}.\n\n"
|
||||
"Examples:\n"
|
||||
'Task: close window → {{"type":"click","x":0.98,"y":0.01,"reason":"X"}}\n'
|
||||
'Task: open notepad → {{"type":"click","x":0.1,"y":0.98,"reason":"search"}}\n'
|
||||
'Task: type hello → {{"type":"type","text":"hello","reason":"input"}}\n\n'
|
||||
"Rules: Mouse only. Use coords from Screen Elements. ONE JSON action.\n"
|
||||
"Action:"
|
||||
)
|
||||
|
||||
try:
|
||||
|
||||
@@ -14,6 +14,7 @@ from typing import List, Optional
|
||||
from openai import AsyncOpenAI
|
||||
|
||||
from .actions import Action, DONE
|
||||
from .prompt_modules import build_system_prompt
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
@@ -165,6 +166,8 @@ class LLMPlanner:
|
||||
context_hint: str = "",
|
||||
scene_elements: int = 0,
|
||||
last_action_type: str = "",
|
||||
pc_state: str = "",
|
||||
os_type: str = "",
|
||||
) -> Action:
|
||||
"""Send screenshot + task to Vision LLM, return parsed Action.
|
||||
|
||||
@@ -178,6 +181,8 @@ class LLMPlanner:
|
||||
context_hint: Memory context / historical experience text.
|
||||
scene_elements: Number of OCR elements detected (for adaptive image).
|
||||
last_action_type: Previous action type (for adaptive image).
|
||||
pc_state: Current PC state (e.g. "desktop", "app_window").
|
||||
os_type: Detected OS type (e.g. "windows", "gnome").
|
||||
"""
|
||||
image_b64 = b64encode(screenshot_bytes).decode("ascii")
|
||||
|
||||
@@ -232,11 +237,17 @@ class LLMPlanner:
|
||||
self.model, len(screenshot_bytes), step + 1, max_steps)
|
||||
if scene_text:
|
||||
logger.debug("Scene text for LLM:\n%s", scene_text[:500])
|
||||
# Use dynamic prompt when state info available, else fallback
|
||||
if pc_state or os_type:
|
||||
system_prompt = build_system_prompt(task, pc_state, os_type)
|
||||
else:
|
||||
system_prompt = AGENT_SYSTEM_PROMPT
|
||||
|
||||
try:
|
||||
response = await self.client.chat.completions.create(
|
||||
model=self.model,
|
||||
messages=[
|
||||
{"role": "system", "content": AGENT_SYSTEM_PROMPT},
|
||||
{"role": "system", "content": system_prompt},
|
||||
{"role": "user", "content": user_content},
|
||||
],
|
||||
max_tokens=self.max_tokens,
|
||||
|
||||
@@ -7,6 +7,8 @@ Calls the centralized NPU Daemon (port 8004) for:
|
||||
|
||||
import json
|
||||
import logging
|
||||
import os
|
||||
import urllib.parse
|
||||
from dataclasses import dataclass
|
||||
from typing import Optional
|
||||
|
||||
@@ -14,6 +16,8 @@ import httpx
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
_UDS_PATH = "/run/npu-daemon/npu.sock"
|
||||
|
||||
|
||||
@dataclass
|
||||
class RedactResult:
|
||||
@@ -27,18 +31,26 @@ class RedactResult:
|
||||
|
||||
|
||||
class NpuClient:
|
||||
"""HTTP client for NPU Daemon."""
|
||||
"""HTTP client for NPU Daemon. Prefers Unix Domain Socket if available."""
|
||||
|
||||
def __init__(self, base_url: str = "http://localhost:8004"):
|
||||
self._base_url = base_url.rstrip("/")
|
||||
self._available: Optional[bool] = None
|
||||
self._use_uds = os.path.exists(_UDS_PATH)
|
||||
|
||||
def _make_client(self, timeout: float = 15.0) -> httpx.AsyncClient:
|
||||
"""Create an httpx client, preferring UDS when available."""
|
||||
if self._use_uds:
|
||||
transport = httpx.AsyncHTTPTransport(uds=_UDS_PATH)
|
||||
return httpx.AsyncClient(transport=transport, timeout=timeout)
|
||||
return httpx.AsyncClient(timeout=timeout)
|
||||
|
||||
async def is_available(self) -> bool:
|
||||
"""Check if NPU Daemon is reachable (cached)."""
|
||||
if self._available is not None:
|
||||
return self._available
|
||||
try:
|
||||
async with httpx.AsyncClient(timeout=3.0) as client:
|
||||
async with self._make_client(timeout=3.0) as client:
|
||||
resp = await client.get(f"{self._base_url}/api/v1/health")
|
||||
self._available = resp.status_code == 200
|
||||
except Exception:
|
||||
@@ -66,30 +78,50 @@ class NpuClient:
|
||||
Raises:
|
||||
httpx.HTTPError: On network/HTTP errors.
|
||||
"""
|
||||
import base64
|
||||
|
||||
types = redact_types or ["id_card", "phone", "bank_card", "email"]
|
||||
|
||||
async with httpx.AsyncClient(timeout=15.0) as client:
|
||||
async with self._make_client() as client:
|
||||
resp = await client.post(
|
||||
f"{self._base_url}/api/v1/privacy/redact-image",
|
||||
files={"image": ("screenshot.jpg", image_bytes, "image/jpeg")},
|
||||
data={"redact_types": json.dumps(types)},
|
||||
headers={"Accept": "image/jpeg"},
|
||||
)
|
||||
resp.raise_for_status()
|
||||
data = resp.json()
|
||||
|
||||
# Decode base64 image if present
|
||||
redacted_image = None
|
||||
if data.get("redacted_image"):
|
||||
redacted_image = base64.b64decode(data["redacted_image"])
|
||||
ct = resp.headers.get("content-type", "")
|
||||
if ct.startswith("image/jpeg"):
|
||||
# Binary mode: JPEG body + metadata in headers
|
||||
redacted_image = resp.content
|
||||
redacted_text = urllib.parse.unquote(
|
||||
resp.headers.get("x-redact-text", "")
|
||||
)
|
||||
findings_raw = resp.headers.get("x-redact-findings", "[]")
|
||||
try:
|
||||
findings = json.loads(findings_raw)
|
||||
except json.JSONDecodeError:
|
||||
findings = []
|
||||
mode = resp.headers.get("x-redact-mode", "text_only")
|
||||
processing_ms = float(resp.headers.get("x-processing-ms", "0"))
|
||||
else:
|
||||
# JSON fallback (backward compatible)
|
||||
import base64 as b64
|
||||
|
||||
data = resp.json()
|
||||
redacted_image = None
|
||||
if data.get("redacted_image"):
|
||||
redacted_image = b64.b64decode(data["redacted_image"])
|
||||
redacted_text = data.get("redacted_text", "")
|
||||
findings = data.get("findings", [])
|
||||
processing_ms = data.get("processing_ms", 0.0)
|
||||
mode = data.get("mode", "text_only")
|
||||
|
||||
return RedactResult(
|
||||
redacted_image=redacted_image,
|
||||
redacted_text=data.get("redacted_text", ""),
|
||||
findings=data.get("findings", []),
|
||||
processing_ms=data.get("processing_ms", 0.0),
|
||||
mode=data.get("mode", "text_only"),
|
||||
redacted_text=redacted_text,
|
||||
findings=findings,
|
||||
processing_ms=processing_ms,
|
||||
mode=mode,
|
||||
)
|
||||
|
||||
async def ocr_analyze(
|
||||
@@ -101,7 +133,7 @@ class NpuClient:
|
||||
|
||||
Returns: {text, regions, processing_ms}
|
||||
"""
|
||||
async with httpx.AsyncClient(timeout=15.0) as client:
|
||||
async with self._make_client() as client:
|
||||
resp = await client.post(
|
||||
f"{self._base_url}/api/v1/ocr/analyze",
|
||||
files={"image": ("image.jpg", image_bytes, "image/jpeg")},
|
||||
|
||||
@@ -0,0 +1,240 @@
|
||||
"""Modular prompt framework for KVM Agent.
|
||||
|
||||
Provides a compact BASE_PROMPT (~350 tokens) plus 9 scenario modules
|
||||
that are dynamically selected based on task keywords, PC state, and OS type.
|
||||
Total budget per request: ≤ 800 tokens (BASE 350 + modules ≤ 450).
|
||||
"""
|
||||
|
||||
import re
|
||||
from typing import List
|
||||
|
||||
# ── Base prompt (always injected, ~350 tokens) ───────────────────
|
||||
|
||||
BASE_PROMPT = """\
|
||||
You are a Computer Use agent controlling a remote PC via KVM.
|
||||
You see screenshots and OCR-detected UI elements with normalized coordinates.
|
||||
|
||||
## Actions (ONE JSON per turn)
|
||||
- {"type":"click","x":0.5,"y":0.3,"button":0,"reason":"label"}
|
||||
x,y: 0.0-1.0 normalized. button: 0=left, 2=right. reason=visible text.
|
||||
- {"type":"type","text":"hello","reason":"..."}
|
||||
- {"type":"scroll","delta":-3,"reason":"..."} (negative=up)
|
||||
- {"type":"wait","delay":2.0,"reason":"..."}
|
||||
- {"type":"done","reason":"..."}
|
||||
|
||||
## Rules
|
||||
1. MOUSE ONLY. No Ctrl/Alt/Win combos. Click visible buttons.
|
||||
2. Use OCR coordinates from Screen Elements — never guess positions.
|
||||
3. "reason" must match exact visible text label for click retry.
|
||||
4. Element size >= 0.04x0.02 = safe target. < 0.02x0.01 = find larger neighbor.
|
||||
5. Taskbar elements: use search bar to launch apps, not direct click.
|
||||
6. If stuck 3 turns: {"type":"done","reason":"stuck"}.
|
||||
7. NEVER type destructive commands.
|
||||
"""
|
||||
|
||||
# ── Scenario modules (~100 tokens each) ──────────────────────────
|
||||
|
||||
MODULE_WINDOW = """\
|
||||
## Window Management
|
||||
- Close: click X at top-right (~0.98, y of title bar). Find title bar via OCR header.
|
||||
- Maximize: double-click title bar center.
|
||||
- Minimize: click _ button left of X (~0.93, title bar Y).
|
||||
- Switch: click target app button in taskbar.
|
||||
- Snap left/right: drag title bar to screen edge.
|
||||
"""
|
||||
|
||||
MODULE_OFFICE = """\
|
||||
## Office Apps (Word/Excel/PPT/WPS/PDF)
|
||||
- Ribbon/toolbar at top: File/Home/Insert/Layout tabs.
|
||||
- Save: File > Save or toolbar floppy icon. Save As: File > Save As.
|
||||
- WPS may show Chinese: 文件/开始/插入/页面布局.
|
||||
- Excel: cells addressed by column-row, formula bar below ribbon.
|
||||
- PPT: slide panel on left, main canvas center, notes at bottom.
|
||||
"""
|
||||
|
||||
MODULE_EMAIL = """\
|
||||
## Email (Outlook/Thunderbird/Webmail)
|
||||
- Three-pane layout: folders left, message list center, preview right.
|
||||
- New email: click New/新建 button (top-left toolbar area).
|
||||
- Reply/Forward: buttons above message preview.
|
||||
- To/CC/Subject fields at top of compose window.
|
||||
- Send button: typically top-left of compose, or bottom.
|
||||
- Attachments: paperclip icon or drag-drop area.
|
||||
"""
|
||||
|
||||
MODULE_BROWSER = """\
|
||||
## Browser (Chrome/Edge/Firefox)
|
||||
- Address bar: top center, click to focus then type URL.
|
||||
- New tab: click + button right of last tab.
|
||||
- Close tab: click X on the tab, NOT window X.
|
||||
- Back/Forward: arrow buttons top-left toolbar.
|
||||
- Downloads: click ... menu > Downloads, or Ctrl+J replaced by menu click.
|
||||
- Search: type in address bar, click search result.
|
||||
- Form fields: click input, then type. For dropdowns, click to expand.
|
||||
"""
|
||||
|
||||
MODULE_FILE = """\
|
||||
## File Management (Explorer/Finder)
|
||||
- Navigation: address bar at top shows current path. Click breadcrumbs.
|
||||
- File dialog (Open/Save): filename field at bottom, type filter dropdown.
|
||||
- Create folder: right-click empty area > New > Folder.
|
||||
- Rename: right-click file > Rename, or click name twice (slow double-click).
|
||||
- Copy/Move: right-click > Copy/Cut, navigate, right-click > Paste.
|
||||
"""
|
||||
|
||||
MODULE_SYSTEM = """\
|
||||
## System Operations
|
||||
- Launch apps: click taskbar search > type name > click result. Wait 2-3s.
|
||||
- Settings: search "settings" or click gear icon.
|
||||
- Task Manager: right-click taskbar > Task Manager.
|
||||
- Volume: click speaker icon in system tray.
|
||||
- WiFi/Network: click network icon in system tray.
|
||||
"""
|
||||
|
||||
MODULE_INTERRUPT = """\
|
||||
## Interruptions (always handle first)
|
||||
- UAC prompt: click Yes/是 to confirm, No/否 to cancel.
|
||||
- Save dialog: Don't Save/不保存 to discard, Save/保存 to keep.
|
||||
- Error dialog: click OK/确定 to dismiss.
|
||||
- Update prompt: click Later/稍后 or X to dismiss.
|
||||
- Right-click menu appeared unexpectedly: click empty area to dismiss.
|
||||
- "Not responding": wait 3s, then click empty area or close via taskbar.
|
||||
"""
|
||||
|
||||
MODULE_ZH_WINDOWS = """\
|
||||
## Chinese Windows Labels
|
||||
- Menu: 文件=File 编辑=Edit 查看=View 格式=Format 工具=Tools 帮助=Help
|
||||
- Dialog: 确定=OK 取消=Cancel 是=Yes 否=No 应用=Apply 浏览=Browse
|
||||
- File ops: 保存=Save 另存为=SaveAs 打开=Open 新建=New 关闭=Close
|
||||
- Edit: 复制=Copy 粘贴=Paste 撤销=Undo 全选=SelectAll 查找=Find
|
||||
- WPS: 开始=Home 插入=Insert 页面布局=Layout 审阅=Review
|
||||
- IME: agent handles toggle. Just type English normally.
|
||||
"""
|
||||
|
||||
MODULE_LINUX = """\
|
||||
## Linux Desktop
|
||||
- GNOME: Activities top-left to open app launcher. Files = Nautilus.
|
||||
- KDE: Application Launcher bottom-left (or kickoff). Files = Dolphin.
|
||||
- XFCE: Whisker Menu or Applications Menu top-left. Files = Thunar.
|
||||
- Terminal: search "terminal" or right-click desktop > Open Terminal.
|
||||
- App install: Software Center / Package Manager, not command line.
|
||||
- File paths: /home/user/Documents, no drive letters.
|
||||
"""
|
||||
|
||||
# ── Module registry ──────────────────────────────────────────────
|
||||
|
||||
_MODULES = {
|
||||
"window": MODULE_WINDOW,
|
||||
"office": MODULE_OFFICE,
|
||||
"email": MODULE_EMAIL,
|
||||
"browser": MODULE_BROWSER,
|
||||
"file": MODULE_FILE,
|
||||
"system": MODULE_SYSTEM,
|
||||
"interrupt": MODULE_INTERRUPT,
|
||||
"zh_windows": MODULE_ZH_WINDOWS,
|
||||
"linux": MODULE_LINUX,
|
||||
}
|
||||
|
||||
# Task keyword → module mapping (supports Chinese + English)
|
||||
_TASK_MODULE_MAP = {
|
||||
r"word|wps|document|文档|docx|excel|表格|xlsx|ppt|演示|pdf": ["office", "file"],
|
||||
r"mail|邮件|outlook|thunderbird|gmail|回复|转发|附件": ["email", "file"],
|
||||
r"browser|浏览器|chrome|edge|firefox|url|网页|搜索|下载|登录": ["browser"],
|
||||
r"file|文件|folder|文件夹|explorer|保存|打开|重命名|删除|复制|移动": ["file"],
|
||||
r"settings|设置|install|安装|network|音量|任务管理器|wifi|蓝牙": ["system"],
|
||||
r"close|关闭|minimize|最大化|最小化|切换|窗口|window|maximize": ["window"],
|
||||
r"open|打开|launch|启动|notepad|calculator|paint|记事本|计算器|画图": ["system", "window"],
|
||||
}
|
||||
|
||||
|
||||
def _select_prompt_modules(
|
||||
task: str, state: str = "", os_type: str = ""
|
||||
) -> List[str]:
|
||||
"""Select relevant prompt modules based on task, state, and OS.
|
||||
|
||||
Args:
|
||||
task: Task description (natural language).
|
||||
state: PC state string (e.g. "desktop", "app_window").
|
||||
os_type: OS type (e.g. "windows", "gnome", "kde").
|
||||
|
||||
Returns:
|
||||
List of module names to include.
|
||||
"""
|
||||
selected = set()
|
||||
|
||||
# Always include interrupt handling
|
||||
selected.add("interrupt")
|
||||
|
||||
# OS-specific module
|
||||
os_lower = os_type.lower() if os_type else ""
|
||||
if os_lower in ("gnome", "kde", "xfce", "linux"):
|
||||
selected.add("linux")
|
||||
elif os_lower in ("windows", "unknown", ""):
|
||||
selected.add("zh_windows")
|
||||
|
||||
# Match task keywords
|
||||
task_lower = task.lower()
|
||||
matched = False
|
||||
for pattern, modules in _TASK_MODULE_MAP.items():
|
||||
if re.search(pattern, task_lower):
|
||||
selected.update(modules)
|
||||
matched = True
|
||||
|
||||
# Default fallback if no keyword matched
|
||||
if not matched:
|
||||
selected.add("window")
|
||||
selected.add("system")
|
||||
|
||||
return sorted(selected)
|
||||
|
||||
|
||||
def build_system_prompt(
|
||||
task: str, pc_state: str = "", os_type: str = ""
|
||||
) -> str:
|
||||
"""Build the complete system prompt with dynamically selected modules.
|
||||
|
||||
Args:
|
||||
task: Task description.
|
||||
pc_state: Current PC state from screen_state detection.
|
||||
os_type: Detected OS type.
|
||||
|
||||
Returns:
|
||||
Complete system prompt string.
|
||||
"""
|
||||
modules = _select_prompt_modules(task, pc_state, os_type)
|
||||
parts = [BASE_PROMPT]
|
||||
for mod_name in modules:
|
||||
if mod_name in _MODULES:
|
||||
parts.append(_MODULES[mod_name])
|
||||
return "\n".join(parts)
|
||||
|
||||
|
||||
def build_context_header(
|
||||
os_type: str = "", state: str = "", window_title: str = "",
|
||||
ime_mode: str = "",
|
||||
) -> str:
|
||||
"""Build compact context header (~15 tokens vs ~40 tokens verbose).
|
||||
|
||||
Example: [OS:WIN] [STATE:APP] [WIN:Notepad - test.txt] [IME:EN]
|
||||
"""
|
||||
parts = []
|
||||
if os_type:
|
||||
os_short = {
|
||||
"windows": "WIN", "gnome": "GNOME", "kde": "KDE",
|
||||
"xfce": "XFCE", "unknown": "?",
|
||||
}.get(os_type.lower(), os_type.upper()[:5])
|
||||
parts.append(f"[OS:{os_short}]")
|
||||
if state:
|
||||
state_short = {
|
||||
"desktop": "DESK", "app_window": "APP", "lock_screen": "LOCK",
|
||||
"sleep": "SLEEP", "bios": "BIOS", "boot": "BOOT",
|
||||
"unknown": "?",
|
||||
}.get(state.lower(), state.upper()[:5])
|
||||
parts.append(f"[STATE:{state_short}]")
|
||||
if window_title:
|
||||
# Truncate long titles
|
||||
title = window_title[:40]
|
||||
parts.append(f"[WIN:{title}]")
|
||||
if ime_mode:
|
||||
parts.append(f"[IME:{ime_mode.upper()[:2]}]")
|
||||
return " ".join(parts)
|
||||
@@ -1,5 +1,5 @@
|
||||
httpx>=0.25.0
|
||||
openai>=1.12.0
|
||||
pyyaml>=6.0
|
||||
aiohttp>=3.8.0
|
||||
pymysql>=1.0.0
|
||||
httpx==0.27.2
|
||||
openai==1.58.1
|
||||
pyyaml==6.0.2
|
||||
aiohttp==3.11.11
|
||||
pymysql==1.1.1
|
||||
|
||||
@@ -28,6 +28,7 @@ class OSType(Enum):
|
||||
WINDOWS = "windows"
|
||||
GNOME = "gnome"
|
||||
KDE = "kde"
|
||||
XFCE = "xfce"
|
||||
UNKNOWN = "unknown"
|
||||
|
||||
|
||||
@@ -62,8 +63,13 @@ DESKTOP_INDICATORS = ["搜索", "开始", "任务栏", "回收站"]
|
||||
# ── OS detection keywords ──────────────────────────────────────
|
||||
WINDOWS_INDICATORS = ["搜索", "search", "开始", "cortana", "edge", "file explorer",
|
||||
"任务栏", "taskbar", "windows"]
|
||||
GNOME_INDICATORS = ["activities", "活动", "nautilus", "gnome", "ubuntu"]
|
||||
KDE_INDICATORS = ["应用程序启动器", "application launcher", "dolphin", "plasma", "kde"]
|
||||
GNOME_INDICATORS = ["activities", "活动", "files", "nautilus", "gnome", "ubuntu",
|
||||
"top bar", "dash to dock"]
|
||||
KDE_INDICATORS = ["plasma", "kickoff", "dolphin", "konsole", "应用程序启动器",
|
||||
"application launcher", "kde"]
|
||||
XFCE_INDICATORS = ["whisker", "thunar", "mousepad", "xfce", "panel"]
|
||||
LINUX_DESKTOP_INDICATORS = ["desktop", "桌面", "file manager", "文件管理器",
|
||||
"terminal", "终端"]
|
||||
|
||||
APP_INDICATORS = [
|
||||
"notepad", "记事本", "powershell", "calculator", "计算器",
|
||||
@@ -129,11 +135,12 @@ class ScreenStateDetector:
|
||||
PCState.LOCK_SCREEN, 0.80, f"keywords={lock_matches}"
|
||||
)
|
||||
|
||||
# 5. Desktop vs app window
|
||||
# 5. Desktop vs app window (includes Linux desktop indicators)
|
||||
has_desktop = any(k.lower() in text_lower for k in DESKTOP_INDICATORS)
|
||||
has_linux_desktop = any(k.lower() in text_lower for k in LINUX_DESKTOP_INDICATORS)
|
||||
has_app = any(k.lower() in text_lower for k in APP_INDICATORS)
|
||||
|
||||
if has_desktop and not has_app:
|
||||
if (has_desktop or has_linux_desktop) and not has_app:
|
||||
return StateDetection(PCState.DESKTOP, 0.75, "desktop_indicators")
|
||||
if has_app:
|
||||
return StateDetection(PCState.APP_WINDOW, 0.70, "app_indicators")
|
||||
@@ -152,6 +159,7 @@ class ScreenStateDetector:
|
||||
OSType.WINDOWS: sum(1 for k in WINDOWS_INDICATORS if k.lower() in text_lower),
|
||||
OSType.GNOME: sum(1 for k in GNOME_INDICATORS if k.lower() in text_lower),
|
||||
OSType.KDE: sum(1 for k in KDE_INDICATORS if k.lower() in text_lower),
|
||||
OSType.XFCE: sum(1 for k in XFCE_INDICATORS if k.lower() in text_lower),
|
||||
}
|
||||
|
||||
best = max(scores, key=scores.get)
|
||||
|
||||
@@ -42,6 +42,7 @@ class TestHybridPlanner:
|
||||
llm = MockLLMPlanner()
|
||||
store = MockTemplateStore(template=None)
|
||||
hp = HybridPlanner(llm, store)
|
||||
hp._local_llm_available = False # Skip local RKLLM path
|
||||
|
||||
result = await hp.plan(b"\x00" * 100, "unknown task", 0, [])
|
||||
assert result.source == "remote_llm"
|
||||
@@ -95,6 +96,7 @@ class TestHybridPlanner:
|
||||
)
|
||||
store = MockTemplateStore(template=template)
|
||||
hp = HybridPlanner(llm, store)
|
||||
hp._local_llm_available = False # Skip local RKLLM path
|
||||
|
||||
# Step 0: template
|
||||
r0 = await hp.plan(b"\x00" * 100, "short task", 0, [])
|
||||
|
||||
Generated
+27
@@ -976,6 +976,15 @@ dependencies = [
|
||||
"pxfm",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "mpp-jpeg"
|
||||
version = "0.1.0"
|
||||
dependencies = [
|
||||
"anyhow",
|
||||
"bindgen",
|
||||
"cc",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "multer"
|
||||
version = "3.1.0"
|
||||
@@ -1030,9 +1039,11 @@ dependencies = [
|
||||
"clap",
|
||||
"face-engine",
|
||||
"image",
|
||||
"mpp-jpeg",
|
||||
"ocr-engine",
|
||||
"regex",
|
||||
"reqwest",
|
||||
"rga",
|
||||
"rknn",
|
||||
"serde",
|
||||
"serde_json",
|
||||
@@ -1067,6 +1078,7 @@ name = "ocr-engine"
|
||||
version = "0.1.0"
|
||||
dependencies = [
|
||||
"image",
|
||||
"rga",
|
||||
"rknn",
|
||||
"serde",
|
||||
"thiserror",
|
||||
@@ -1318,6 +1330,21 @@ dependencies = [
|
||||
"web-sys",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "rga"
|
||||
version = "0.1.0"
|
||||
dependencies = [
|
||||
"anyhow",
|
||||
"rga-sys",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "rga-sys"
|
||||
version = "0.1.0"
|
||||
dependencies = [
|
||||
"bindgen",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "ring"
|
||||
version = "0.17.14"
|
||||
|
||||
@@ -2,6 +2,9 @@
|
||||
members = [
|
||||
"crates/rknn-sys",
|
||||
"crates/rknn",
|
||||
"crates/rga-sys",
|
||||
"crates/rga",
|
||||
"crates/mpp-jpeg",
|
||||
"crates/ocr-engine",
|
||||
"crates/face-engine",
|
||||
".",
|
||||
@@ -20,6 +23,8 @@ edition.workspace = true
|
||||
|
||||
[dependencies]
|
||||
rknn = { path = "crates/rknn" }
|
||||
rga = { path = "crates/rga" }
|
||||
mpp-jpeg = { path = "crates/mpp-jpeg" }
|
||||
ocr-engine = { path = "crates/ocr-engine" }
|
||||
face-engine = { path = "crates/face-engine" }
|
||||
|
||||
|
||||
@@ -0,0 +1,12 @@
|
||||
[package]
|
||||
name = "mpp-jpeg"
|
||||
version.workspace = true
|
||||
edition.workspace = true
|
||||
build = "build.rs"
|
||||
|
||||
[build-dependencies]
|
||||
bindgen = "0.71"
|
||||
cc = "1"
|
||||
|
||||
[dependencies]
|
||||
anyhow = "1"
|
||||
@@ -0,0 +1,62 @@
|
||||
use std::env;
|
||||
use std::path::PathBuf;
|
||||
|
||||
fn main() {
|
||||
println!("cargo:rustc-link-lib=dylib=rockchip_mpp");
|
||||
println!("cargo:rustc-link-search=native=/usr/lib/aarch64-linux-gnu");
|
||||
println!("cargo:rerun-if-changed=wrapper.h");
|
||||
println!("cargo:rerun-if-changed=wrapper.c");
|
||||
|
||||
// Compile wrapper.c — provides real C functions for MPP buffer macros
|
||||
cc::Build::new()
|
||||
.file("wrapper.c")
|
||||
.include("/usr/include")
|
||||
.compile("mpp_jpeg_wrapper");
|
||||
|
||||
let bindings = bindgen::Builder::default()
|
||||
.header("wrapper.h")
|
||||
.clang_arg("-I/usr/include")
|
||||
// Core API functions
|
||||
.allowlist_function("mpp_create")
|
||||
.allowlist_function("mpp_init")
|
||||
.allowlist_function("mpp_destroy")
|
||||
.allowlist_function("mpp_enc_cfg_init")
|
||||
.allowlist_function("mpp_enc_cfg_deinit")
|
||||
.allowlist_function("mpp_enc_cfg_set_s32")
|
||||
.allowlist_function("mpp_enc_cfg_set_u32")
|
||||
.allowlist_function("mpp_frame_init")
|
||||
.allowlist_function("mpp_frame_deinit")
|
||||
.allowlist_function("mpp_frame_set_.*")
|
||||
.allowlist_function("mpp_packet_get_.*")
|
||||
.allowlist_function("mpp_packet_deinit")
|
||||
.allowlist_function("mpp_buffer_group_put")
|
||||
// Wrapper functions for macros (declared in wrapper.h, defined in wrapper.c)
|
||||
.allowlist_function("mpp_jpeg_buffer_get")
|
||||
.allowlist_function("mpp_jpeg_buffer_put")
|
||||
.allowlist_function("mpp_jpeg_buffer_get_ptr")
|
||||
.allowlist_function("mpp_jpeg_buffer_group_get")
|
||||
// Key types
|
||||
.allowlist_type("MppCtx")
|
||||
.allowlist_type("MppApi")
|
||||
.allowlist_type("MppEncCfg")
|
||||
.allowlist_type("MppFrame")
|
||||
.allowlist_type("MppPacket")
|
||||
.allowlist_type("MppBuffer")
|
||||
.allowlist_type("MppBufferGroup")
|
||||
.allowlist_type("MppCodingType")
|
||||
.allowlist_type("MppCtxType")
|
||||
.allowlist_type("MppFrameFormat")
|
||||
.allowlist_type("MppEncRcMode_e")
|
||||
.allowlist_type("MppBufferType")
|
||||
.allowlist_type("MPP_RET")
|
||||
.derive_debug(true)
|
||||
.derive_default(true)
|
||||
.parse_callbacks(Box::new(bindgen::CargoCallbacks::new()))
|
||||
.generate()
|
||||
.expect("Unable to generate MPP JPEG bindings");
|
||||
|
||||
let out_path = PathBuf::from(env::var("OUT_DIR").unwrap());
|
||||
bindings
|
||||
.write_to_file(out_path.join("bindings.rs"))
|
||||
.expect("Couldn't write bindings");
|
||||
}
|
||||
@@ -0,0 +1,232 @@
|
||||
//! MPP Hardware JPEG Encoder for Rust.
|
||||
//!
|
||||
//! Uses Rockchip VEPU to encode NV12/RGB → JPEG in hardware (~5ms per 1080p frame).
|
||||
//! Independent MPP context — does not interfere with any other MPP encoder/decoder.
|
||||
|
||||
#![allow(non_upper_case_globals)]
|
||||
#![allow(non_camel_case_types)]
|
||||
#![allow(non_snake_case)]
|
||||
|
||||
mod sys {
|
||||
include!(concat!(env!("OUT_DIR"), "/bindings.rs"));
|
||||
}
|
||||
|
||||
use anyhow::{bail, Result};
|
||||
use std::ptr;
|
||||
|
||||
// Re-export bindgen-prefixed constants with short names
|
||||
const MPP_OK: sys::MPP_RET = sys::MPP_RET_MPP_OK;
|
||||
const MPP_CTX_ENC: sys::MppCtxType = sys::MppCtxType_MPP_CTX_ENC;
|
||||
const MPP_VIDEO_CodingMJPEG: sys::MppCodingType = sys::MppCodingType_MPP_VIDEO_CodingMJPEG;
|
||||
const MPP_FMT_YUV420SP: sys::MppFrameFormat = sys::MppFrameFormat_MPP_FMT_YUV420SP;
|
||||
const MPP_ENC_RC_MODE_FIXQP: sys::MppEncRcMode_e = sys::MppEncRcMode_e_MPP_ENC_RC_MODE_FIXQP;
|
||||
const MPP_ENC_GET_CFG: sys::MpiCmd = sys::MpiCmd_MPP_ENC_GET_CFG;
|
||||
const MPP_ENC_SET_CFG: sys::MpiCmd = sys::MpiCmd_MPP_ENC_SET_CFG;
|
||||
const MPP_BUFFER_TYPE_ION: sys::MppBufferType = sys::MppBufferType_MPP_BUFFER_TYPE_ION;
|
||||
|
||||
/// Hardware JPEG encoder using Rockchip MPP VEPU.
|
||||
pub struct MppJpegEncoder {
|
||||
ctx: sys::MppCtx,
|
||||
mpi: *mut sys::MppApi,
|
||||
cfg: sys::MppEncCfg,
|
||||
buf_grp: sys::MppBufferGroup,
|
||||
width: u32,
|
||||
height: u32,
|
||||
}
|
||||
|
||||
// SAFETY: MppCtx is a thread-safe handle (MPP serializes internally)
|
||||
unsafe impl Send for MppJpegEncoder {}
|
||||
unsafe impl Sync for MppJpegEncoder {}
|
||||
|
||||
impl MppJpegEncoder {
|
||||
/// Create a new JPEG encoder for the given resolution.
|
||||
/// quality: 1-99 (higher = better quality, larger file).
|
||||
pub fn new(width: u32, height: u32, quality: u32) -> Result<Self> {
|
||||
let quality = quality.clamp(1, 99);
|
||||
|
||||
let mut ctx: sys::MppCtx = ptr::null_mut();
|
||||
let mut mpi: *mut sys::MppApi = ptr::null_mut();
|
||||
|
||||
// Create MPP context
|
||||
let ret = unsafe { sys::mpp_create(&mut ctx, &mut mpi) };
|
||||
if ret != MPP_OK {
|
||||
bail!("mpp_create failed: {}", ret);
|
||||
}
|
||||
|
||||
// Initialize as MJPEG encoder
|
||||
let ret = unsafe { sys::mpp_init(ctx, MPP_CTX_ENC as _, MPP_VIDEO_CodingMJPEG as _) };
|
||||
if ret != MPP_OK {
|
||||
unsafe { sys::mpp_destroy(ctx) };
|
||||
bail!("mpp_init MJPEG failed: {}", ret);
|
||||
}
|
||||
|
||||
// Create config
|
||||
let mut cfg: sys::MppEncCfg = ptr::null_mut();
|
||||
let ret = unsafe { sys::mpp_enc_cfg_init(&mut cfg) };
|
||||
if ret != MPP_OK {
|
||||
unsafe { sys::mpp_destroy(ctx) };
|
||||
bail!("mpp_enc_cfg_init failed: {}", ret);
|
||||
}
|
||||
|
||||
// Get current config
|
||||
let ret = unsafe { (*mpi).control.unwrap()(ctx, MPP_ENC_GET_CFG as _, cfg as _) };
|
||||
if ret != MPP_OK {
|
||||
unsafe {
|
||||
sys::mpp_enc_cfg_deinit(cfg);
|
||||
sys::mpp_destroy(ctx);
|
||||
}
|
||||
bail!("MPP_ENC_GET_CFG failed: {}", ret);
|
||||
}
|
||||
|
||||
// Configure encoder
|
||||
unsafe {
|
||||
let w = width as i32;
|
||||
let h = height as i32;
|
||||
|
||||
sys::mpp_enc_cfg_set_s32(cfg, b"prep:width\0".as_ptr() as _, w);
|
||||
sys::mpp_enc_cfg_set_s32(cfg, b"prep:height\0".as_ptr() as _, h);
|
||||
sys::mpp_enc_cfg_set_s32(cfg, b"prep:hor_stride\0".as_ptr() as _, w);
|
||||
sys::mpp_enc_cfg_set_s32(cfg, b"prep:ver_stride\0".as_ptr() as _, h);
|
||||
sys::mpp_enc_cfg_set_s32(
|
||||
cfg,
|
||||
b"prep:format\0".as_ptr() as _,
|
||||
MPP_FMT_YUV420SP as i32,
|
||||
);
|
||||
|
||||
sys::mpp_enc_cfg_set_s32(
|
||||
cfg,
|
||||
b"codec:type\0".as_ptr() as _,
|
||||
MPP_VIDEO_CodingMJPEG as i32,
|
||||
);
|
||||
sys::mpp_enc_cfg_set_s32(cfg, b"jpeg:q_factor\0".as_ptr() as _, quality as i32);
|
||||
|
||||
sys::mpp_enc_cfg_set_s32(
|
||||
cfg,
|
||||
b"rc:mode\0".as_ptr() as _,
|
||||
MPP_ENC_RC_MODE_FIXQP as i32,
|
||||
);
|
||||
sys::mpp_enc_cfg_set_s32(cfg, b"rc:fps_in_num\0".as_ptr() as _, 1);
|
||||
sys::mpp_enc_cfg_set_s32(cfg, b"rc:fps_in_denorm\0".as_ptr() as _, 1);
|
||||
sys::mpp_enc_cfg_set_s32(cfg, b"rc:fps_out_num\0".as_ptr() as _, 1);
|
||||
sys::mpp_enc_cfg_set_s32(cfg, b"rc:fps_out_denorm\0".as_ptr() as _, 1);
|
||||
}
|
||||
|
||||
// Apply config
|
||||
let ret = unsafe { (*mpi).control.unwrap()(ctx, MPP_ENC_SET_CFG as _, cfg as _) };
|
||||
if ret != MPP_OK {
|
||||
unsafe {
|
||||
sys::mpp_enc_cfg_deinit(cfg);
|
||||
sys::mpp_destroy(ctx);
|
||||
}
|
||||
bail!("MPP_ENC_SET_CFG failed: {}", ret);
|
||||
}
|
||||
|
||||
// Buffer group
|
||||
let mut buf_grp: sys::MppBufferGroup = ptr::null_mut();
|
||||
let ret = unsafe { sys::mpp_jpeg_buffer_group_get(&mut buf_grp, MPP_BUFFER_TYPE_ION) };
|
||||
if ret != MPP_OK {
|
||||
unsafe {
|
||||
sys::mpp_enc_cfg_deinit(cfg);
|
||||
sys::mpp_destroy(ctx);
|
||||
}
|
||||
bail!("mpp_buffer_group_get failed: {}", ret);
|
||||
}
|
||||
|
||||
Ok(MppJpegEncoder {
|
||||
ctx,
|
||||
mpi,
|
||||
cfg,
|
||||
buf_grp,
|
||||
width,
|
||||
height,
|
||||
})
|
||||
}
|
||||
|
||||
/// Encode NV12 data to JPEG. Returns the JPEG bytes.
|
||||
pub fn encode_nv12(&self, nv12_data: &[u8]) -> Result<Vec<u8>> {
|
||||
let frame_size = (self.width * self.height * 3 / 2) as usize;
|
||||
if nv12_data.len() < frame_size {
|
||||
bail!(
|
||||
"NV12 buffer too small: {} < {}",
|
||||
nv12_data.len(),
|
||||
frame_size
|
||||
);
|
||||
}
|
||||
|
||||
unsafe {
|
||||
// Allocate MPP buffer
|
||||
let mut buffer: sys::MppBuffer = ptr::null_mut();
|
||||
let ret = sys::mpp_jpeg_buffer_get(self.buf_grp, &mut buffer, frame_size);
|
||||
if ret != MPP_OK {
|
||||
bail!("mpp_buffer_get failed: {}", ret);
|
||||
}
|
||||
|
||||
// Copy NV12 data to MPP buffer
|
||||
let buf_ptr = sys::mpp_jpeg_buffer_get_ptr(buffer);
|
||||
ptr::copy_nonoverlapping(nv12_data.as_ptr(), buf_ptr as *mut u8, frame_size);
|
||||
|
||||
// Create frame
|
||||
let mut frame: sys::MppFrame = ptr::null_mut();
|
||||
sys::mpp_frame_init(&mut frame);
|
||||
sys::mpp_frame_set_width(frame, self.width);
|
||||
sys::mpp_frame_set_height(frame, self.height);
|
||||
sys::mpp_frame_set_hor_stride(frame, self.width);
|
||||
sys::mpp_frame_set_ver_stride(frame, self.height);
|
||||
sys::mpp_frame_set_fmt(frame, MPP_FMT_YUV420SP as _);
|
||||
sys::mpp_frame_set_buffer(frame, buffer);
|
||||
|
||||
// Encode
|
||||
let ret = (*self.mpi).encode_put_frame.unwrap()(self.ctx, frame);
|
||||
if ret != MPP_OK {
|
||||
sys::mpp_jpeg_buffer_put(buffer);
|
||||
sys::mpp_frame_deinit(&mut frame);
|
||||
bail!("encode_put_frame failed: {}", ret);
|
||||
}
|
||||
|
||||
let mut packet: sys::MppPacket = ptr::null_mut();
|
||||
let ret = (*self.mpi).encode_get_packet.unwrap()(self.ctx, &mut packet);
|
||||
if ret != MPP_OK {
|
||||
sys::mpp_jpeg_buffer_put(buffer);
|
||||
sys::mpp_frame_deinit(&mut frame);
|
||||
bail!("encode_get_packet failed: {}", ret);
|
||||
}
|
||||
|
||||
let result = if !packet.is_null() {
|
||||
let pkt_data = sys::mpp_packet_get_data(packet);
|
||||
let pkt_size = sys::mpp_packet_get_length(packet);
|
||||
|
||||
if !pkt_data.is_null() && pkt_size > 0 {
|
||||
let jpeg =
|
||||
std::slice::from_raw_parts(pkt_data as *const u8, pkt_size).to_vec();
|
||||
sys::mpp_packet_deinit(&mut packet);
|
||||
Ok(jpeg)
|
||||
} else {
|
||||
sys::mpp_packet_deinit(&mut packet);
|
||||
bail!("MPP returned empty JPEG packet")
|
||||
}
|
||||
} else {
|
||||
bail!("MPP returned null packet")
|
||||
};
|
||||
|
||||
sys::mpp_jpeg_buffer_put(buffer);
|
||||
sys::mpp_frame_deinit(&mut frame);
|
||||
result
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
impl Drop for MppJpegEncoder {
|
||||
fn drop(&mut self) {
|
||||
unsafe {
|
||||
if !self.buf_grp.is_null() {
|
||||
sys::mpp_buffer_group_put(self.buf_grp);
|
||||
}
|
||||
if !self.cfg.is_null() {
|
||||
sys::mpp_enc_cfg_deinit(self.cfg);
|
||||
}
|
||||
if !self.ctx.is_null() {
|
||||
sys::mpp_destroy(self.ctx);
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,21 @@
|
||||
// Wrapper implementations for MPP buffer macros.
|
||||
// These macros expand to _with_tag/_with_caller variants with __FUNCTION__ etc.
|
||||
// bindgen cannot handle this, so we provide real C functions.
|
||||
|
||||
#include <rockchip/mpp_buffer.h>
|
||||
|
||||
MPP_RET mpp_jpeg_buffer_get(MppBufferGroup group, MppBuffer *buffer, size_t size) {
|
||||
return mpp_buffer_get_with_tag(group, buffer, size, "mpp-jpeg", __FUNCTION__);
|
||||
}
|
||||
|
||||
MPP_RET mpp_jpeg_buffer_put(MppBuffer buffer) {
|
||||
return mpp_buffer_put_with_caller(buffer, __FUNCTION__);
|
||||
}
|
||||
|
||||
void *mpp_jpeg_buffer_get_ptr(MppBuffer buffer) {
|
||||
return mpp_buffer_get_ptr_with_caller(buffer, __FUNCTION__);
|
||||
}
|
||||
|
||||
MPP_RET mpp_jpeg_buffer_group_get(MppBufferGroup *group, MppBufferType type) {
|
||||
return mpp_buffer_group_get(group, type, MPP_BUFFER_INTERNAL, "mpp-jpeg", __FUNCTION__);
|
||||
}
|
||||
@@ -0,0 +1,12 @@
|
||||
// MPP JPEG encoder API — subset of Rockchip MPP needed for JPEG encoding
|
||||
#include <rockchip/rk_mpi.h>
|
||||
#include <rockchip/mpp_frame.h>
|
||||
#include <rockchip/mpp_packet.h>
|
||||
#include <rockchip/mpp_buffer.h>
|
||||
|
||||
// Wrapper functions for MPP macros (bindgen can't expand C macros).
|
||||
// Implemented in wrapper.c, linked via cc crate.
|
||||
MPP_RET mpp_jpeg_buffer_get(MppBufferGroup group, MppBuffer *buffer, size_t size);
|
||||
MPP_RET mpp_jpeg_buffer_put(MppBuffer buffer);
|
||||
void *mpp_jpeg_buffer_get_ptr(MppBuffer buffer);
|
||||
MPP_RET mpp_jpeg_buffer_group_get(MppBufferGroup *group, MppBufferType type);
|
||||
@@ -5,6 +5,7 @@ edition.workspace = true
|
||||
|
||||
[dependencies]
|
||||
rknn = { path = "../rknn" }
|
||||
rga = { path = "../rga" }
|
||||
image = { version = "0.25", default-features = false, features = ["jpeg", "png"] }
|
||||
tracing = "0.1"
|
||||
thiserror = "2"
|
||||
|
||||
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user