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:
2026-05-17 14:50:50 +00:00
co-authored by Claude Opus 4.7
parent 0e2c765ea7
commit 9bb799b26a
139 changed files with 3497 additions and 5289 deletions
+6
View File
@@ -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/
+47 -24
View File
@@ -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)
```
## 编码规范
- Pythonasyncio + httpx,类型注解,dataclass 优先
- 异步优先:所有 I/O 操作使用 async/await
- Rusttokio + axum + reqwestserde 序列化,thiserror 错误处理
- Python(仅 kvm_agent,待迁移):asyncio + httpx,类型注解,dataclass 优先
- 异步优先:所有 I/O 操作使用 async/await (Rust: tokio, Python: asyncio)
- 错误处理:finally 块中释放资源(HID 控制、隐私模式)
- 测试:pytest + pytest-asynciomock 外部依赖
- 配置:YAML 文件 + 环境变量dataclass 承载
- 测试:Rust #[tokio::test]Python pytest + pytest-asynciomock 外部依赖
- 配置: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 每秒触发 oneshotnative 模式不需要 | 从包中移除 |
## 目标设备
- 硬件:NanoPC-T6 (RK3588, 8核 ARM64, 6 TOPS NPU)
- 硬件:NanoPC-T6 (RK3588, 8核 ARM64, 8GB RAM, 6 TOPS NPU)
- 开发与运行同机:代码编辑、编译、服务运行均在本机完成,资源共享
- 系统:Debian/Ubuntu ARM64
- Python3.12.3
- Go1.22+
- NPUrknn-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 TabAuditPage 2 Tab,新增 5 条 API 路由 |
+1 -3
View File
@@ -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
+2 -5
View File
@@ -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
View File
@@ -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
-3
View File
@@ -1,3 +0,0 @@
#!/bin/bash
export PYTHONPATH=/usr/lib/kvm-agent
exec python3 -m kvm_agent "$@"
+1
View File
@@ -0,0 +1 @@
/usr/sbin/kvm-agent
+1
View File
@@ -0,0 +1 @@
/etc/kvm-bridge/bridge.env
+6 -7
View File
@@ -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).
+13 -14
View File
@@ -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
View File
@@ -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
+7 -9
View File
@@ -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
View File
@@ -0,0 +1,2 @@
BRIDGE_DATA_DIR=/var/lib/kvm-bridge
EMBED_MODEL_PATH=/usr/share/kvm-bridge/models/embedder.onnx
+28
View File
@@ -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
View File
@@ -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
View File
@@ -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. knowledgelong_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.faissIndexFlatIP)。"""
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.faissmmap=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_termsummary + 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 # 有序 dictPython 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 ServiceOpenAI 兼容代理 + 复杂度路由 + 记忆注入。"""
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:
# 无后端时创建禁用状态的 QueryRewriterenabled=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]
-27
View File
@@ -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
View File
@@ -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()
+6 -5
View File
@@ -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).
+7 -7
View File
@@ -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.
+11 -15
View File
@@ -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
+3 -14
View File
@@ -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
-2
View File
@@ -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
View File
@@ -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
View File
@@ -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
-3
View File
@@ -1,3 +0,0 @@
#!/bin/bash
export PYTHONPATH=/usr/lib/kvm-mitm
exec python3 -m privacy_gateway.mitm_launcher "$@"
+1
View File
@@ -0,0 +1 @@
/etc/npu-daemon/config.yaml
+1
View File
@@ -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.
Vendored Executable
+8
View File
@@ -0,0 +1,8 @@
#!/bin/bash
set -e
case "$1" in
purge)
rm -rf /etc/npu-daemon
;;
esac
exit 0
+18 -10
View File
@@ -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
View File
@@ -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
+2
View File
@@ -0,0 +1,2 @@
/etc/kvm-privacy/pii_rules.yaml
/etc/kvm-privacy/surnames.txt
+9
View File
@@ -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
View File
@@ -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:
+3 -2
View File
@@ -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 APINPU_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)
# ── 后端 CNPU 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
# ── 后端 Amediapipe-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
# ── 后端 Bmediapipe_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 APINPU_DAEMON_URL 环境变量)
2. ppocr_rknn.PPOcrRknnrknnlite,板端友好)
3. ppocr_det + ppocr_recrknn 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
# ── 后端 CNPU 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
# ── 后端 Appocr_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
# ── 后端 Bppocr_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()
+3 -4
View File
@@ -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.
+17 -1
View File
@@ -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
View File
@@ -0,0 +1,8 @@
#!/bin/bash
set -e
case "$1" in
purge)
rm -rf /var/lib/kvm-rkllm
;;
esac
exit 0
-3
View File
@@ -1,3 +0,0 @@
#!/bin/bash
export PYTHONPATH=/usr/lib/kvm-rkllm
exec python3 -m rkllm_server "$@"
+3 -2
View File
@@ -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
+9 -7
View File
@@ -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
-27
View File
@@ -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
-18
View File
@@ -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
-26
View File
@@ -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
+5 -6
View File
@@ -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
-21
View File
@@ -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
+2
View File
@@ -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:
+5 -5
View File
@@ -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
View File
@@ -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"
+8 -17
View File
@@ -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(
+747 -9
View File
File diff suppressed because it is too large Load Diff
+13
View File
@@ -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
View File
@@ -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
+18 -7
View File
@@ -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:
+12 -1
View File
@@ -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,
+47 -15
View File
@@ -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")},
+240
View File
@@ -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)
+5 -5
View File
@@ -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
+12 -4
View File
@@ -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, [])
+27
View File
@@ -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"
+5
View File
@@ -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