wangxi 3 týždňov pred
rodič
commit
83e45b14d9

+ 4 - 2
.env.example

@@ -1,7 +1,9 @@
-NEO4J_URI=neo4j://127.0.0.1:7687
+NEO4J_URI=http://121.43.55.7:7474
 NEO4J_USER=neo4j
 NEO4J_PASSWORD=change-me
+NEO4J_DATABASE=neo4j
 DEEPSEEK_API_KEY=sk-xxxx
 DEEPSEEK_BASE_URL=https://api.deepseek.com
 DEEPSEEK_MODEL=deepseek-v4-flash
-
+REDIS_URL=redis://127.0.0.1:6379/0
+REDIS_CONTEXT_TTL_SECONDS=1296000

+ 55 - 0
AGENTS.md

@@ -0,0 +1,55 @@
+# 开始规则
+在开始进行代码撰写前,请先完成以下内容:
+- 完整阅读当前文件内容,它定义了整个项目的主体架构
+- 阅读docs/技术方案.md,该文件包含了需要实现的功能
+- 运行 bash init.sh 验证项目环境是否正常
+- 阅读 feature_list.json 查看最近的一些工作内容
+- 阅读 READEME.md 了解重要功能文件的说明
+
+# 项目整体结构
+目前该项目主要分为以下部分:
+
+- 请求日志记录(api_log)
+
+该文件夹主要记录了目前服务接口的请求参数、agent处理流程以及最终的结果,还有请求响应时间
+- 参考数据(data)
+
+里面包含了大量混杂的实际生产数据(data/20260529_申勤提供数据、data/慧管理遗产、data/投标文件资料_20260625、data/收到的数据),以及由生产数据规整出来的后续用于构建知识图谱的数据模板(data/templates),以及模拟出来的测试知识图谱数据(data/test_data)
+data/config_example.json 为构建知识图谱所需的
+- 文档文件(docs)
+
+里面包含了整体功能说明、以及目前的技术实施路线(docs/技术方案.md)、处理技术流程图(docs/images)以及dms数据库和目前模板数据之间的映射关系(初步映射关系图,目前可忽略该映射关系,后续有最新的映射关系文件output/模板_源数据_DMS字段关联.csv)
+
+- 前端代码(html)
+
+里面包含了目前整个项目的前端代码
+
+- 编码模型(models)
+
+里面存放了一个编码模型,用于计算用户问题和知识图谱实体之间的相似度,以便于用户问题和知识图谱之间进行关联
+
+- dms与知识图谱数据模板的关联(output)
+
+里面存放了目前dms数据库与知识图谱模板字段之间的关联关系,目前最新关联关系文件为output/模板_源数据_DMS字段关联.csv。以及知识图谱的元知识图谱数据output/meta_graph_schema.json
+
+- 测试文件(scripts)
+
+一些编程智能体在开发过程中开发的测试文件
+
+- 主要功能代码(src)
+
+包含了根据知识图谱数据进行知识图谱构建、问答agent构建的功能代码
+
+- 配置文件(.env)
+
+整体项目的一些配置信息,比如neo4j数据库的连接信息、问答智能体基座模型的一些配置信息、redis配置信息
+
+- 服务运行环境配置信息
+
+该项目使用uv进行项目环境控制,相关信息分别存放在:.venv、pyproject.toml、uv.lock中
+
+- 项目主题主题功能介绍(README.md)
+
+里面包含了目前各部分功能的主要应用代码说明以及
+
+

+ 11 - 2
README.md

@@ -12,6 +12,7 @@
 - **多轮对话**:同一会话内保留结构化上下文(问题/实体主键/回答),支持拆解式提问与指代;
   多轮答案复用(`reuse_check`)由 LLM 结合 当前问题+历史问题+历史回答 判断是否可直接推导;
 - **API 服务**:FastAPI + 全异步 `astream`,`thread_id` 会话管理,SSE 流式处理进度,多请求并发;
+- **上下文存储**:会话状态由 Redis 异步 checkpointer 持久化,避免进程内存在长期增长;
 - **人工确认**:槽位确认、计划确认、子图范围收缩,均可自动化或走 human-in-the-loop;
 - **只读保障**:所有问答 LLM 提示词注入只读约束,问答阶段禁止修改图谱数据。
 
@@ -24,6 +25,7 @@
 | LLM | DeepSeek(`deepseek-v4-flash`,OpenAI 兼容) |
 | 语义嵌入 | Qwen3-Embedding-0.6B(本地模型) |
 | Agent 编排 | LangGraph |
+| 上下文存储 | Redis(自实现 `PlainRedisSaver`,原生 Redis 即可,无需模块) |
 | API | FastAPI + Uvicorn(全异步 astream) |
 
 ## 快速开始
@@ -35,7 +37,7 @@ uv sync
 cp .env.example .env
 ```
 
-编辑 `.env`:
+编辑 `.env`(实际配置值统一维护在这里;`config.py` 只做类型化读取)
 
 ```ini
 NEO4J_URI=neo4j://127.0.0.1:7687
@@ -44,12 +46,16 @@ NEO4J_PASSWORD=你的密码
 DEEPSEEK_API_KEY=sk-你的key
 DEEPSEEK_BASE_URL=https://api.deepseek.com
 DEEPSEEK_MODEL=deepseek-v4-flash
+REDIS_URL=redis://127.0.0.1:6379/0
+REDIS_CONTEXT_TTL_SECONDS=1296000  # 上下文保留 15 天
 ```
 
 ### 2. 启动 Neo4j
 
 在 Neo4j Desktop(或服务)中启动数据库,确认 Bolt 端口(默认 7687)可连接。
 
+启动原生 Redis(无需 RedisJSON / RediSearch 模块),确认 `REDIS_URL` 可连接。
+
 ### 3. 准备嵌入模型
 
 将 Qwen3-Embedding-0.6B 放到 `models/Qwen3-Embedding-0.6B`(API 服务启动时会预加载;CLI 首次问答会加载,之后进程内缓存)。
@@ -115,6 +121,9 @@ uv run python -m knowledge_agent.api --host 0.0.0.0  # 局域网访问(需放
 默认监听 `127.0.0.1:8000`(`--host 0.0.0.0` 时局域网内通过 `http://<你的IP>:8000` 访问),
 交互文档见 `http://127.0.0.1:8000/docs`。
 
+服务启动时会连接 `REDIS_URL` 并创建原生 Redis checkpointer;若 Redis 未启动,启动会失败。
+`REDIS_CONTEXT_TTL_SECONDS` 控制上下文在 Redis 中的保留时间,当前默认配置为 15 天(1296000 秒)。
+
 > ⚠️ 当前未接鉴权,`0.0.0.0` 对外开放前请先加 API Key / 权限校验。
 
 ### 接口
@@ -271,7 +280,7 @@ knowledge_agent/
 ├── models/Qwen3-Embedding-0.6B
 ├── scripts/                 # 模板/测试数据/构建/问答 CLI/DMS 映射等
 └── src/knowledge_agent/
-    ├── config.py            # .env 配置
+    ├── config.py            # .env 类型化读取(不内置默认值)
     ├── db.py                # Neo4j driver 单例
     ├── meta/                # 元知识图谱(单一事实源)
     ├── graph/               # 图谱构建(schemas/reader/builder)

+ 8 - 6
docs/技术方案.md

@@ -419,7 +419,8 @@ Cypher 与参数(`QueryResult.query/query_params`),供 `--debug` 输出。
   (只看某项目 / 按岗位筛选 / 汇总口径 / 继续返回全部),选择后回到计划节点重查;
   同一问题只询问一次(`scope_offered` 防死循环);
 - **上下文记忆**:LangGraph checkpointer(thread_id)+ messages 历史,支持多轮
-  (如"改成4月""换成税务局");槽位抽取会带入最近对话上下文;
+  (如"改成4月""换成税务局");槽位抽取会带入最近对话上下文;API 服务使用
+  Redis 异步 checkpointer 持久化,CLI 仍使用进程内 InMemorySaver;
 - **上下文复用**:多轮中当前问题若可由历史轮次直接推导(如上一轮已列出各月考勤,
   这轮问"连续三个月都有加班的人"),`reuse_check` 判定可复用后由 `answer_reuse`
   直接回答(标注"基于前序对话推导"),不重复检索,省 2~3 次 LLM 调用与全部查询;
@@ -448,17 +449,18 @@ Cypher 与参数(`QueryResult.query/query_params`),供 `--debug` 输出。
   预留 `on_chunk` 回调,后续可升级为 token 级透传;
 - **全异步 astream**:/ask 与 /ask/stream 共用 `_astream_run` 事件流(节点级进度 + interrupt 循环),
   resume 走 `_astream_resume`,同一套逻辑,无 invoke/stream 双实现;
-- **并发**:全局共享一个 compiled graph + `InMemorySaver`,状态按 `thread_id` 隔离,
+- **并发**:全局共享一个 compiled graph + `PlainRedisSaver`,状态按 `thread_id` 隔离,
   FastAPI async + `graph.astream` 天然支持多请求并发;
-- **上下文记忆**:仍为内存(`InMemorySaver`),后续持久化/图数据库化时替换
-  `api.py` 的 checkpointer 初始化与 `ask.py` 的 `_make_checkpointer()` 即可。
+- **上下文记忆**:API 已改为 Redis 持久化(自实现 `PlainRedisSaver`,原生 Redis 即可),进程重启不丢上下文;
+  上下文按 `REDIS_CONTEXT_TTL_SECONDS` 设置过期时间(默认 15 天);
+  CLI 的 `scripts/ask.py` 仍使用 `InMemorySaver`,仅用于本地调试。
 
 ## 7. 代码结构
 
 ```text
 knowledge_agent/
 ├── pyproject.toml / uv.lock
-├── .env                     # NEO4J_* / DEEPSEEK_*(不入库)
+├── .env                     # NEO4J_* / DEEPSEEK_* / REDIS_URL / TTL(不入库)
 ├── docs/技术方案.md
 ├── data/
 │   ├── templates/           # 15 类单表模板(+填写说明)
@@ -555,7 +557,7 @@ knowledge_agent/
    - 机制:回答只依据检索子图事实,缺失明确说"数据中未找到",金额带单位、人名/项目用标准名称。
 3. **凭据管理**
    - 位置:`.env`(gitignore 排除)+ `src/knowledge_agent/config.py`;
-   - 机制:Neo4j / DeepSeek 连接信息从环境变量读取;`.env.example` 仅含占位值。
+   - 机制:Neo4j / DeepSeek / Redis 连接信息只维护在 `.env`,`config.py` 仅做类型化读取;`.env.example` 仅含占位值。
 
 ### 11.4 权限模型(🟡 数据就绪,待接线)
 

+ 1 - 0
pyproject.toml

@@ -16,6 +16,7 @@ dependencies = [
     "networkx>=3.3",
     "openai<3",
     "langgraph>=1.2.11",
+    "redis>=5.2.1",
     "langchain-openai>=1.4.3",
     "sentence-transformers>=5.7.0",
     "pycryptodome>=3.23.0",

+ 303 - 0
src/knowledge_agent/agent/redis_checkpointer.py

@@ -0,0 +1,303 @@
+"""极简 Redis 异步 checkpoint saver(不依赖 RediSearch / RedisJSON)。
+
+只使用 Redis String、Hash 与 SCAN,适配原生 Redis 5+。
+"""
+from __future__ import annotations
+
+from collections.abc import AsyncIterator, Sequence
+from types import TracebackType
+from typing import Any
+from urllib.parse import quote, unquote
+
+from langchain_core.runnables import RunnableConfig
+from langgraph.checkpoint.base import (
+    WRITES_IDX_MAP,
+    BaseCheckpointSaver,
+    ChannelVersions,
+    Checkpoint,
+    CheckpointMetadata,
+    CheckpointTuple,
+    get_checkpoint_id,
+    get_checkpoint_metadata,
+)
+from redis.asyncio import Redis
+
+
+class PlainRedisSaver(BaseCheckpointSaver[str]):
+    """原生 Redis 异步 checkpoint saver。
+
+    每个 checkpoint 存为一个 Redis Hash,字段为 type/data(JsonPlusSerializer 的
+    序列化结果),latest 指针存为普通 String。这样不要求 Redis 安装 RedisJSON 或
+    RediSearch 模块。
+    """
+
+    def __init__(
+        self,
+        redis_url: str,
+        *,
+        key_prefix: str = "ka",
+        ttl_seconds: int | None = None,
+        serde=None,
+    ) -> None:
+        super().__init__(serde=serde)
+        self.redis_url = redis_url
+        self.key_prefix = key_prefix
+        self.ttl_seconds = ttl_seconds
+        self._redis: Redis | None = None
+
+    @property
+    def redis(self) -> Redis:
+        if self._redis is None:
+            raise RuntimeError("Redis client 尚未初始化")
+        return self._redis
+
+    async def asetup(self) -> "PlainRedisSaver":
+        self._redis = Redis.from_url(self.redis_url, decode_responses=False)
+        await self._redis.ping()
+        return self
+
+    async def __aenter__(self) -> "PlainRedisSaver":
+        await self.asetup()
+        return self
+
+    async def __aexit__(
+        self,
+        exc_type: type[BaseException] | None,
+        exc: BaseException | None,
+        tb: TracebackType | None,
+    ) -> None:
+        if self._redis is not None:
+            await self._redis.aclose()
+            self._redis = None
+
+    @staticmethod
+    def _part(value: str | int) -> str:
+        return quote(str(value), safe="")
+
+    def _threads_key(self) -> str:
+        return f"{self.key_prefix}:threads"
+
+    def _cp_key(self, thread_id: str, checkpoint_ns: str, checkpoint_id: str) -> str:
+        return (
+            f"{self.key_prefix}:cp:"
+            f"{self._part(thread_id)}:{self._part(checkpoint_ns)}:"
+            f"{self._part(checkpoint_id)}"
+        )
+
+    def _writes_key(self, thread_id: str, checkpoint_ns: str, checkpoint_id: str) -> str:
+        return (
+            f"{self.key_prefix}:writes:"
+            f"{self._part(thread_id)}:{self._part(checkpoint_ns)}:"
+            f"{self._part(checkpoint_id)}"
+        )
+
+    def _latest_key(self, thread_id: str, checkpoint_ns: str) -> str:
+        return f"{self.key_prefix}:latest:{self._part(thread_id)}:{self._part(checkpoint_ns)}"
+
+    async def _save_obj(self, key: str, obj: Any) -> None:
+        type_, data = self.serde.dumps_typed(obj)
+        await self.redis.hset(key, mapping={"type": type_, "data": data})
+        if self.ttl_seconds is not None:
+            await self.redis.expire(key, self.ttl_seconds)
+
+    async def _load_obj(self, key: str) -> Any:
+        pipe = self.redis.pipeline()
+        pipe.hget(key, "type")
+        pipe.hget(key, "data")
+        type_, data = await pipe.execute()
+        if type_ is None or data is None:
+            return None
+        return self.serde.loads_typed((type_.decode("utf-8"), data))
+
+    async def _get_writes(
+        self, thread_id: str, checkpoint_ns: str, checkpoint_id: str
+    ) -> list[tuple[str, str, Any]]:
+        stored = await self._load_obj(self._writes_key(thread_id, checkpoint_ns, checkpoint_id))
+        if not isinstance(stored, list):
+            return []
+        pending: list[tuple[str, str, Any]] = []
+        for row in stored:
+            if not isinstance(row, (list, tuple)) or len(row) != 4:
+                continue
+            task_id, channel, _idx, value = row
+            pending.append((task_id, channel, value))
+        return pending
+
+    async def aget_tuple(self, config: RunnableConfig) -> CheckpointTuple | None:
+        thread_id: str = config["configurable"]["thread_id"]
+        checkpoint_ns: str = config["configurable"].get("checkpoint_ns", "")
+        requested_id = get_checkpoint_id(config)
+        checkpoint_id = requested_id
+
+        if not checkpoint_id:
+            raw_latest = await self.redis.get(self._latest_key(thread_id, checkpoint_ns))
+            checkpoint_id = raw_latest.decode("utf-8") if raw_latest else None
+
+        if not checkpoint_id:
+            return None
+
+        payload = await self._load_obj(self._cp_key(thread_id, checkpoint_ns, checkpoint_id))
+        if not isinstance(payload, dict):
+            return None
+
+        checkpoint = payload.get("checkpoint")
+        metadata = payload.get("metadata", {})
+        parent_checkpoint_id = payload.get("parent_checkpoint_id")
+        pending_writes = await self._get_writes(thread_id, checkpoint_ns, checkpoint_id)
+
+        if requested_id:
+            return_config: RunnableConfig = config
+        else:
+            return_config = {
+                "configurable": {
+                    "thread_id": thread_id,
+                    "checkpoint_ns": checkpoint_ns,
+                    "checkpoint_id": checkpoint_id,
+                }
+            }
+
+        parent_config: RunnableConfig | None = None
+        if parent_checkpoint_id:
+            parent_config = {
+                "configurable": {
+                    "thread_id": thread_id,
+                    "checkpoint_ns": checkpoint_ns,
+                    "checkpoint_id": parent_checkpoint_id,
+                }
+            }
+
+        return CheckpointTuple(
+            config=return_config,
+            checkpoint=checkpoint,
+            metadata=metadata,
+            parent_config=parent_config,
+            pending_writes=pending_writes,
+        )
+
+    async def aput(
+        self,
+        config: RunnableConfig,
+        checkpoint: Checkpoint,
+        metadata: CheckpointMetadata,
+        new_versions: ChannelVersions,
+    ) -> RunnableConfig:
+        thread_id: str = config["configurable"]["thread_id"]
+        checkpoint_ns: str = config["configurable"].get("checkpoint_ns", "")
+        config_checkpoint_id = config["configurable"].get("checkpoint_id")
+        checkpoint_id = checkpoint.get("id") or config_checkpoint_id
+
+        if not checkpoint_id:
+            raise RuntimeError("checkpoint 缺少 checkpoint_id")
+
+        parent_checkpoint_id = None
+        if config_checkpoint_id and config_checkpoint_id != checkpoint_id:
+            parent_checkpoint_id = config_checkpoint_id
+
+        payload = {
+            "checkpoint": checkpoint,
+            "metadata": get_checkpoint_metadata(config, metadata),
+            "parent_checkpoint_id": parent_checkpoint_id,
+        }
+        await self._save_obj(self._cp_key(thread_id, checkpoint_ns, checkpoint_id), payload)
+        if self.ttl_seconds is not None:
+            await self.redis.set(
+                self._latest_key(thread_id, checkpoint_ns), checkpoint_id, ex=self.ttl_seconds
+            )
+        else:
+            await self.redis.set(self._latest_key(thread_id, checkpoint_ns), checkpoint_id)
+        await self.redis.sadd(self._threads_key(), thread_id)
+
+        return {
+            "configurable": {
+                "thread_id": thread_id,
+                "checkpoint_ns": checkpoint_ns,
+                "checkpoint_id": checkpoint_id,
+            }
+        }
+
+    async def aput_writes(
+        self,
+        config: RunnableConfig,
+        writes: Sequence[tuple[str, Any]],
+        task_id: str,
+        task_path: str = "",
+    ) -> None:
+        thread_id: str = config["configurable"]["thread_id"]
+        checkpoint_ns: str = config["configurable"].get("checkpoint_ns", "")
+        checkpoint_id = config["configurable"].get("checkpoint_id")
+        if not checkpoint_id:
+            raw_latest = await self.redis.get(self._latest_key(thread_id, checkpoint_ns))
+            checkpoint_id = raw_latest.decode("utf-8") if raw_latest else None
+        if not checkpoint_id:
+            return
+        key = self._writes_key(thread_id, checkpoint_ns, checkpoint_id)
+
+        stored = await self._load_obj(key)
+        records: list[list[Any]] = [
+            list(row) for row in stored if isinstance(row, (list, tuple)) and len(row) == 4
+        ] if isinstance(stored, list) else []
+
+        for idx, (channel, value) in enumerate(writes):
+            write_idx = WRITES_IDX_MAP.get(channel, idx)
+            if write_idx >= 0 and any(
+                row[0] == task_id and row[2] == write_idx for row in records
+            ):
+                continue
+            records.append([task_id, channel, write_idx, value])
+
+        await self._save_obj(key, records)
+
+    async def alist(
+        self,
+        config: RunnableConfig | None,
+        *,
+        filter: dict[str, Any] | None = None,
+        before: RunnableConfig | None = None,
+        limit: int | None = None,
+    ) -> AsyncIterator[CheckpointTuple]:
+        if config is None:
+            return
+        thread_id: str = config["configurable"]["thread_id"]
+        checkpoint_ns: str = config["configurable"].get("checkpoint_ns", "")
+
+        match = f"{self.key_prefix}:cp:{self._part(thread_id)}:{self._part(checkpoint_ns)}:*"
+        keys = [key async for key in self.redis.scan_iter(match=match)]
+        checkpoint_ids: list[str] = []
+        for key in keys:
+            try:
+                checkpoint_ids.append(unquote(key.decode("utf-8").rsplit(":", 1)[-1]))
+            except Exception:
+                continue
+
+        count = 0
+        for checkpoint_id in sorted(checkpoint_ids, reverse=True):
+            tup = await self.aget_tuple(
+                {
+                    "configurable": {
+                        "thread_id": thread_id,
+                        "checkpoint_ns": checkpoint_ns,
+                        "checkpoint_id": checkpoint_id,
+                    }
+                }
+            )
+            if tup is None:
+                continue
+            yield tup
+            count += 1
+            if limit is not None and count >= limit:
+                return
+
+    async def adelete_thread(self, thread_id: str) -> None:
+        tid = self._part(thread_id)
+        patterns = (
+            f"{self.key_prefix}:cp:{tid}:*",
+            f"{self.key_prefix}:writes:{tid}:*",
+            f"{self.key_prefix}:latest:{tid}:*",
+        )
+        keys: list[bytes] = []
+        for pattern in patterns:
+            async for key in self.redis.scan_iter(match=pattern):
+                keys.append(key)
+        if keys:
+            await self.redis.delete(*keys)
+        await self.redis.srem(self._threads_key(), thread_id)

+ 32 - 15
src/knowledge_agent/api.py

@@ -28,24 +28,35 @@ from fastapi import FastAPI
 from fastapi.middleware.cors import CORSMiddleware
 from fastapi.responses import FileResponse, StreamingResponse
 from fastapi.staticfiles import StaticFiles
-from langgraph.checkpoint.memory import InMemorySaver
+from .agent.redis_checkpointer import PlainRedisSaver
 from langgraph.types import Command
 from pydantic import BaseModel
 
 from .agent.graph import build_agent_graph
+from .config import get_settings
 
 
 @asynccontextmanager
 async def lifespan(_: FastAPI):
+    global _graph
+
     from .agent.embedding import preload_model
 
-    try:
-        preload_model()
-        print("Qwen3-Embedding-0.6B 模型已在服务启动时预加载", flush=True)
-    except Exception as exc:  # noqa: BLE001
-        # 模型加载失败不应阻塞服务启动;请求时会再次尝试并返回更完整错误。
-        print(f"Qwen3 模型预加载失败,将在首次请求时重试:{exc}", flush=True)
-    yield
+    settings = get_settings()
+    async with PlainRedisSaver(
+        settings.redis_url, ttl_seconds=settings.redis_context_ttl_seconds
+    ) as checkpointer:
+        _graph = build_agent_graph(checkpointer=checkpointer)
+        print("Redis 异步 checkpointer 已就绪", flush=True)
+
+        try:
+            preload_model()
+            print("Qwen3-Embedding-0.6B 模型已在服务启动时预加载", flush=True)
+        except Exception as exc:  # noqa: BLE001
+            # 模型加载失败不应阻塞服务启动;请求时会再次尝试并返回更完整错误。
+            print("Qwen3 模型预加载失败,将在首次请求时重试:" + str(exc), flush=True)
+        yield
+    _graph = None
 
 
 app = FastAPI(title="申勤物业知识助手", version="0.1.0", lifespan=lifespan)
@@ -56,8 +67,8 @@ app.add_middleware(
     allow_headers=["*"],
 )
 
-# 全局共享图:状态按 thread_id 隔离(InMemorySaver),多请求并发安全
-_graph = build_agent_graph(checkpointer=InMemorySaver())
+# 全局共享图:状态按 thread_id 隔离,由 PlainRedisSaver 持久化到原生 Redis
+_graph = None
 API_LOG_DIR = Path(__file__).resolve().parents[2] / "api_log"
 HTML_DIR = Path(__file__).resolve().parents[2] / "html"
 OUTPUT_DIR = Path(__file__).resolve().parents[2] / "output"
@@ -104,6 +115,12 @@ def _config(thread_id: str) -> dict:
     return {"configurable": {"thread_id": thread_id}}
 
 
+def _get_graph():
+    if _graph is None:
+        raise RuntimeError("Agent graph 尚未初始化,请通过 FastAPI lifespan 启动服务")
+    return _graph
+
+
 def _safe_thread_id(thread_id: str) -> str:
     return re.sub(r"[^A-Za-z0-9_.-]+", "_", str(thread_id or "default"))[:64] or "default"
 
@@ -176,7 +193,7 @@ async def _astream_run(thread_id: str, query: str,
     inp: object = {"question": query, "user_id": "api", "enable_reuse": enable_reuse}
     while True:
         interrupted: object | None = None
-        async for item in _graph.astream(inp, config, stream_mode=["updates", "custom"]):
+        async for item in _get_graph().astream(inp, config, stream_mode=["updates", "custom"]):
             if isinstance(item, tuple) and len(item) == 2:
                 mode, payload = item
             elif isinstance(item, dict):
@@ -203,7 +220,7 @@ async def _astream_run(thread_id: str, query: str,
                    "options": intr.get("options", [])}
             return
         inp = Command(resume="确认")
-    st = await _graph.aget_state(config)
+    st = await _get_graph().aget_state(config)
     state = st.values or {}
     _write_api_log(endpoint, thread_id, {
         "query": query,
@@ -233,7 +250,7 @@ async def _astream_resume(thread_id: str, reply: str,
     inp: object = Command(resume=reply)
     while True:
         interrupted: object | None = None
-        async for item in _graph.astream(inp, config, stream_mode=["updates", "custom"]):
+        async for item in _get_graph().astream(inp, config, stream_mode=["updates", "custom"]):
             if isinstance(item, tuple) and len(item) == 2:
                 mode, payload = item
             elif isinstance(item, dict):
@@ -258,7 +275,7 @@ async def _astream_resume(thread_id: str, reply: str,
                "message": intr.get("message"),
                "options": intr.get("options", [])}
         return
-    st = await _graph.aget_state(config)
+    st = await _get_graph().aget_state(config)
     state = st.values or {}
     _write_api_log(endpoint, thread_id, {"reply": reply}, state,
                    round(time.monotonic() - t0, 3))
@@ -373,7 +390,7 @@ async def resume(thread_id: str, body: ResumeRequest):
 
 @app.get("/threads/{thread_id}/history")
 async def history(thread_id: str):
-    st = await _graph.aget_state(_config(thread_id))
+    st = await _get_graph().aget_state(_config(thread_id))
     values = st.values or {}
     return {
         "thread_id": thread_id,

+ 28 - 8
src/knowledge_agent/config.py

@@ -1,4 +1,4 @@
-"""项目配置:从 .env 读取 Neo4j / DeepSeek 连接信息。"""
+"""项目配置:统一从 .env 读取;config.py 只做类型化访问,不再内置环境默认值。"""
 
 from __future__ import annotations
 
@@ -15,24 +15,44 @@ def _load() -> None:
     load_dotenv(root / ".env")
 
 
+def _required(name: str) -> str:
+    value = os.getenv(name, "").strip()
+    if not value:
+        raise RuntimeError(f"缺少环境变量 {name},请在 .env 中配置")
+    return value
+
+
+def _required_int(name: str) -> int:
+    raw = _required(name)
+    try:
+        return int(raw)
+    except ValueError as exc:
+        raise RuntimeError(f"环境变量 {name} 必须是整数,当前为 {raw!r}") from exc
+
+
 @dataclass(frozen=True)
 class Settings:
     neo4j_uri: str
     neo4j_user: str
     neo4j_password: str
+    neo4j_database: str
     deepseek_api_key: str
     deepseek_base_url: str
     deepseek_model: str
+    redis_url: str
+    redis_context_ttl_seconds: int
 
 
 def get_settings() -> Settings:
     _load()
     return Settings(
-        neo4j_uri=os.getenv("NEO4J_URI", "neo4j://127.0.0.1:7687"),
-        neo4j_user=os.getenv("NEO4J_USER", "neo4j"),
-        neo4j_password=os.getenv("NEO4J_PASSWORD", ""),
-        deepseek_api_key=os.getenv("DEEPSEEK_API_KEY", ""),
-        deepseek_base_url=os.getenv("DEEPSEEK_BASE_URL", "https://api.deepseek.com"),
-        deepseek_model=os.getenv("DEEPSEEK_MODEL", "deepseek-flash"),
+        neo4j_uri=_required("NEO4J_URI"),
+        neo4j_user=_required("NEO4J_USER"),
+        neo4j_password=_required("NEO4J_PASSWORD"),
+        neo4j_database=_required("NEO4J_DATABASE"),
+        deepseek_api_key=_required("DEEPSEEK_API_KEY"),
+        deepseek_base_url=_required("DEEPSEEK_BASE_URL"),
+        deepseek_model=_required("DEEPSEEK_MODEL"),
+        redis_url=_required("REDIS_URL"),
+        redis_context_ttl_seconds=_required_int("REDIS_CONTEXT_TTL_SECONDS"),
     )
-

+ 91 - 10
src/knowledge_agent/db.py

@@ -1,22 +1,103 @@
-"""Neo4j driver 单例:连接池复用,避免每次查询都新建连接。"""
-
+"""Neo4j HTTP Query API client, compatible with the previous Bolt driver interface."""
 from __future__ import annotations
 
-from neo4j import GraphDatabase
+import base64
+import json
+import urllib.error
+import urllib.request
+from typing import Any
 
 from .config import get_settings
 
 
-_driver = None
+class _QueryResult:
+    def __init__(self, records: list[dict[str, Any]]):
+        self.records = records
+
+
+class _Session:
+    def __init__(self, client: "HttpNeo4jDriver"):
+        self._client = client
+
+    def __enter__(self):
+        return self
+
+    def __exit__(self, exc_type, exc, tb):
+        return False
+
+    def run(self, query: str, **params) -> list[dict[str, Any]]:
+        return self._client._query(query, params)
+
+
+class HttpNeo4jDriver:
+    def __init__(self, uri: str, user: str, password: str, database: str = "neo4j"):
+        self.uri = uri.rstrip("/")
+        self.user = user
+        self.password = password
+        self.database = database
+
+    def __enter__(self):
+        return self
+
+    def __exit__(self, exc_type, exc, tb):
+        return False
+
+    def close(self):
+        return None
 
+    def execute_query(self, query: str, **params) -> _QueryResult:
+        return _QueryResult(self._query(query, params))
 
-def get_driver():
+    def session(self) -> _Session:
+        return _Session(self)
+
+    def _query(self, query: str, params: dict[str, Any]) -> list[dict[str, Any]]:
+        url = f"{self.uri}/db/{self.database}/query/v2"
+        payload = {"statement": query, "parameters": params}
+        token = base64.b64encode(f"{self.user}:{self.password}".encode()).decode()
+        req = urllib.request.Request(
+            url,
+            data=json.dumps(payload, ensure_ascii=False).encode("utf-8"),
+            headers={
+                "Content-Type": "application/json",
+                "Accept": "application/json",
+                "Authorization": f"Basic {token}",
+            },
+            method="POST",
+        )
+        try:
+            with urllib.request.urlopen(req, timeout=30) as resp:
+                body = resp.read().decode("utf-8")
+        except urllib.error.HTTPError as exc:
+            detail = exc.read().decode("utf-8", errors="replace")
+            raise RuntimeError(f"Neo4j HTTP {exc.code}: {detail}") from exc
+        except urllib.error.URLError as exc:
+            raise RuntimeError(f"Neo4j HTTP connection failed: {exc.reason}") from exc
+
+        data = json.loads(body)
+        errors = data.get("errors") or []
+        if errors:
+            raise RuntimeError(f"Neo4j query failed: {errors}")
+        result_data = data.get("data") or {}
+        fields = result_data.get("fields") or []
+        values = result_data.get("values") or []
+        return [dict(zip(fields, row)) for row in values]
+
+
+_driver: HttpNeo4jDriver | None = None
+
+
+def get_driver(uri: str | None = None, user: str | None = None, password: str | None = None) -> HttpNeo4jDriver:
     global _driver
-    if _driver is None:
+    if uri or user or password:
         s = get_settings()
-        _driver = GraphDatabase.driver(
-            s.neo4j_uri,
-            auth=(s.neo4j_user, s.neo4j_password),
-            connection_timeout=3.0,   # 连接失败快速返回,避免启动探测长时间重试
+        return HttpNeo4jDriver(
+            uri or s.neo4j_uri,
+            user or s.neo4j_user,
+            password or s.neo4j_password,
+            database=s.neo4j_database,
         )
+    if _driver is None:
+        s = get_settings()
+        _driver = HttpNeo4jDriver(s.neo4j_uri, s.neo4j_user, s.neo4j_password, database=s.neo4j_database)
     return _driver

+ 0 - 1
src/knowledge_agent/graph/builder.py

@@ -24,7 +24,6 @@ def create_constraints(driver) -> None:
         "CREATE CONSTRAINT IF NOT EXISTS FOR (n:项目) REQUIRE n.编号 IS UNIQUE",
         "CREATE CONSTRAINT IF NOT EXISTS FOR (n:人员) REQUIRE n.工号 IS UNIQUE",
         "CREATE CONSTRAINT IF NOT EXISTS FOR (n:片区) REQUIRE n.编号 IS UNIQUE",
-        "CREATE CONSTRAINT IF NOT EXISTS FOR (n:设备) REQUIRE (n.项目编号, n.设备编号) IS NODE KEY",
     ):
         _run(driver, stmt)
 

+ 2 - 2
src/knowledge_agent/graph/pipeline.py

@@ -4,7 +4,7 @@ from __future__ import annotations
 
 from typing import Any
 
-from neo4j import GraphDatabase
+from ..db import get_driver
 
 from ..config import get_settings
 from .builder import build_all, clear_graph, create_constraints, graph_summary
@@ -32,7 +32,7 @@ def build_knowledge_graph(
     if not any(records.values()):
         return {"ok": False, "records": {}, "issues": [str(i) for i in issues], "graph": {}}
 
-    with GraphDatabase.driver(uri, auth=(user, password)) as driver:
+    with get_driver(uri, user, password) as driver:
         if clear_first:
             clear_graph(driver)
         create_constraints(driver)

+ 4 - 4
src/knowledge_agent/meta/builder.py

@@ -4,7 +4,7 @@ from __future__ import annotations
 
 import json
 
-from neo4j import GraphDatabase
+from ..db import get_driver
 
 from ..config import get_settings
 from .profiler import build_report
@@ -116,9 +116,9 @@ def build(db) -> dict:
 
 def main() -> None:
     s = get_settings()
-    with GraphDatabase.driver(s.neo4j_uri, auth=(s.neo4j_user, s.neo4j_password)) as driver:
-        _clear(driver)
-        result = build(driver)
+    driver = get_driver()
+    _clear(driver)
+    result = build(driver)
     print("元知识图谱构建完成:")
     for k, v in result.items():
         print(f"  {k}: {v}")

+ 41 - 9
uv.lock

@@ -5,11 +5,14 @@ resolution-markers = [
     "python_full_version >= '3.14' and sys_platform == 'win32'",
     "python_full_version >= '3.14' and sys_platform == 'emscripten'",
     "python_full_version >= '3.14' and sys_platform != 'emscripten' and sys_platform != 'win32'",
-    "python_full_version >= '3.12' and python_full_version < '3.14' and sys_platform == 'win32'",
+    "python_full_version == '3.13.*' and sys_platform == 'win32'",
+    "python_full_version == '3.12.*' and sys_platform == 'win32'",
     "python_full_version < '3.12' and sys_platform == 'win32'",
-    "python_full_version >= '3.12' and python_full_version < '3.14' and sys_platform == 'emscripten'",
+    "python_full_version == '3.13.*' and sys_platform == 'emscripten'",
+    "python_full_version == '3.12.*' and sys_platform == 'emscripten'",
     "python_full_version < '3.12' and sys_platform == 'emscripten'",
-    "python_full_version >= '3.12' and python_full_version < '3.14' and sys_platform != 'emscripten' and sys_platform != 'win32'",
+    "python_full_version == '3.13.*' and sys_platform != 'emscripten' and sys_platform != 'win32'",
+    "python_full_version == '3.12.*' and sys_platform != 'emscripten' and sys_platform != 'win32'",
     "python_full_version < '3.12' and sys_platform != 'emscripten' and sys_platform != 'win32'",
 ]
 
@@ -44,6 +47,15 @@ wheels = [
     { url = "https://files.pythonhosted.org/packages/da/35/f2287558c17e29fafc8ef3daf819bb9834061cfa43bff8014f7df7f63bdc/anyio-4.14.2-py3-none-any.whl", hash = "sha256:9f505dda5ac9f0c8309b5e8bd445a8c2bf7246f3ce950121e45ea15bc41d1494", size = 125813, upload-time = "2026-07-12T20:29:05.763Z" },
 ]
 
+[[package]]
+name = "async-timeout"
+version = "5.0.1"
+source = { registry = "https://pypi.org/simple" }
+sdist = { url = "https://files.pythonhosted.org/packages/a5/ae/136395dfbfe00dfc94da3f3e136d0b13f394cba8f4841120e34226265780/async_timeout-5.0.1.tar.gz", hash = "sha256:d9321a7a3d5a6a5e187e824d2fa0793ce379a202935782d555d6e9d2735677d3", size = 9274, upload-time = "2024-11-06T16:41:39.6Z" }
+wheels = [
+    { url = "https://files.pythonhosted.org/packages/fe/ba/e2081de779ca30d473f21f5b30e0e737c438205440784c7dfc81efc2b029/async_timeout-5.0.1-py3-none-any.whl", hash = "sha256:39e3809566ff85354557ec2398b55e096c8364bacac9405a7a1fa429e77fe76c", size = 6233, upload-time = "2024-11-06T16:41:37.9Z" },
+]
+
 [[package]]
 name = "certifi"
 version = "2026.7.22"
@@ -804,6 +816,7 @@ dependencies = [
     { name = "pypdf" },
     { name = "python-docx" },
     { name = "python-dotenv" },
+    { name = "redis" },
     { name = "sentence-transformers" },
     { name = "uvicorn", extra = ["standard"] },
     { name = "xlrd" },
@@ -825,6 +838,7 @@ requires-dist = [
     { name = "pypdf", specifier = ">=4.0" },
     { name = "python-docx", specifier = ">=1.1" },
     { name = "python-dotenv", specifier = ">=1.0" },
+    { name = "redis", specifier = ">=5.2.1" },
     { name = "sentence-transformers", specifier = ">=5.7.0" },
     { name = "uvicorn", extras = ["standard"], specifier = ">=0.52.2" },
     { name = "xlrd", specifier = ">=2.0" },
@@ -1353,9 +1367,12 @@ resolution-markers = [
     "python_full_version >= '3.14' and sys_platform == 'win32'",
     "python_full_version >= '3.14' and sys_platform == 'emscripten'",
     "python_full_version >= '3.14' and sys_platform != 'emscripten' and sys_platform != 'win32'",
-    "python_full_version >= '3.12' and python_full_version < '3.14' and sys_platform == 'win32'",
-    "python_full_version >= '3.12' and python_full_version < '3.14' and sys_platform == 'emscripten'",
-    "python_full_version >= '3.12' and python_full_version < '3.14' and sys_platform != 'emscripten' and sys_platform != 'win32'",
+    "python_full_version == '3.13.*' and sys_platform == 'win32'",
+    "python_full_version == '3.12.*' and sys_platform == 'win32'",
+    "python_full_version == '3.13.*' and sys_platform == 'emscripten'",
+    "python_full_version == '3.12.*' and sys_platform == 'emscripten'",
+    "python_full_version == '3.13.*' and sys_platform != 'emscripten' and sys_platform != 'win32'",
+    "python_full_version == '3.12.*' and sys_platform != 'emscripten' and sys_platform != 'win32'",
 ]
 sdist = { url = "https://files.pythonhosted.org/packages/9a/80/db0b4559e57ec36362bedbb05530a87fafbcb6067708c946967a41d449e7/numpy-2.5.2.tar.gz", hash = "sha256:d482d171c406ae88c5b19cad3b6a1c4c5209f886ab74bc44c2c865c23f52d860", size = 20773161, upload-time = "2026-08-09T13:48:27.962Z" }
 wheels = [
@@ -2146,6 +2163,18 @@ wheels = [
     { url = "https://files.pythonhosted.org/packages/f1/12/de94a39c2ef588c7e6455cfbe7343d3b2dc9d6b6b2f40c4c6565744c873d/pyyaml-6.0.3-cp314-cp314t-win_arm64.whl", hash = "sha256:ebc55a14a21cb14062aa4162f906cd962b28e2e9ea38f9b4391244cd8de4ae0b", size = 149341, upload-time = "2025-09-25T21:32:56.828Z" },
 ]
 
+[[package]]
+name = "redis"
+version = "7.4.1"
+source = { registry = "https://pypi.org/simple" }
+dependencies = [
+    { name = "async-timeout", marker = "python_full_version < '3.11.3'" },
+]
+sdist = { url = "https://files.pythonhosted.org/packages/51/93/05e7d4a65285066a74f48697f9b9cde5cfce71398033d69ed83c3d98f5c9/redis-7.4.1.tar.gz", hash = "sha256:1a1df5067062cf7cbe677994e391f8ee0840f499d370f1a71266e0dd3aa9308e", size = 4945742, upload-time = "2026-06-05T09:10:06.703Z" }
+wheels = [
+    { url = "https://files.pythonhosted.org/packages/4a/2e/2677f3f93dae0497e7e33b6637302e7f3744efc553f34231183e32584885/redis-7.4.1-py3-none-any.whl", hash = "sha256:1fa4647af1c5e93a2c685aa248ee44cce092691146d41390518dabe9a99839b0", size = 410171, upload-time = "2026-06-05T09:10:05.128Z" },
+]
+
 [[package]]
 name = "regex"
 version = "2026.7.19"
@@ -2445,9 +2474,12 @@ resolution-markers = [
     "python_full_version >= '3.14' and sys_platform == 'win32'",
     "python_full_version >= '3.14' and sys_platform == 'emscripten'",
     "python_full_version >= '3.14' and sys_platform != 'emscripten' and sys_platform != 'win32'",
-    "python_full_version >= '3.12' and python_full_version < '3.14' and sys_platform == 'win32'",
-    "python_full_version >= '3.12' and python_full_version < '3.14' and sys_platform == 'emscripten'",
-    "python_full_version >= '3.12' and python_full_version < '3.14' and sys_platform != 'emscripten' and sys_platform != 'win32'",
+    "python_full_version == '3.13.*' and sys_platform == 'win32'",
+    "python_full_version == '3.12.*' and sys_platform == 'win32'",
+    "python_full_version == '3.13.*' and sys_platform == 'emscripten'",
+    "python_full_version == '3.12.*' and sys_platform == 'emscripten'",
+    "python_full_version == '3.13.*' and sys_platform != 'emscripten' and sys_platform != 'win32'",
+    "python_full_version == '3.12.*' and sys_platform != 'emscripten' and sys_platform != 'win32'",
 ]
 dependencies = [
     { name = "numpy", version = "2.5.2", source = { registry = "https://pypi.org/simple" } },