feat: initial commit

This commit is contained in:
zkcoi committed 2026-06-23 20:29:02 +08:00
commit ea8f6066c4
217 files changed
+60754

No files matched your search

+106
View File
@@ -0,0 +1,106 @@
#!/usr/bin/env python3
# -*- coding: utf-8 -*-
"""
Data Modules - 数据链模块包。
注意:
- 这里采用延迟导入(lazy import),避免在执行 `python -m data_modules.xxx` 时,
因包级 __init__ 提前导入子模块而触发 runpy 的 RuntimeWarning。
- 推荐用法永远安全:
from data_modules.index_manager import IndexManager
但为了兼容历史代码,也保留:
from data_modules import IndexManager
"""
from __future__ import annotations
from importlib import import_module
from typing import Any
__all__ = [
# Config
"DataModulesConfig",
"get_config",
"set_project_root",
# API Client
"ModalAPIClient",
"get_client",
# Entity Linker
"EntityLinker",
"DisambiguationResult",
# State Manager
"StateManager",
"EntityState",
"Relationship",
"StateChange",
# Index Manager
"IndexManager",
"ChapterMeta",
"SceneMeta",
"ReviewMetrics",
"RelationshipEventMeta",
# RAG Adapter
"RAGAdapter",
"SearchResult",
"ContextManager",
"ContextRanker",
"SnapshotManager",
"QueryRouter",
# Style Sampler
"StyleSampler",
"StyleSample",
"SceneType",
]
_LAZY_EXPORTS: dict[str, tuple[str, str]] = {
# Config
"DataModulesConfig": (".config", "DataModulesConfig"),
"get_config": (".config", "get_config"),
"set_project_root": (".config", "set_project_root"),
# API Client
"ModalAPIClient": (".api_client", "ModalAPIClient"),
"get_client": (".api_client", "get_client"),
# Entity Linker
"EntityLinker": (".entity_linker", "EntityLinker"),
"DisambiguationResult": (".entity_linker", "DisambiguationResult"),
# State Manager
"StateManager": (".state_manager", "StateManager"),
"EntityState": (".state_manager", "EntityState"),
"Relationship": (".state_manager", "Relationship"),
"StateChange": (".state_manager", "StateChange"),
# Index Manager
"IndexManager": (".index_manager", "IndexManager"),
"ChapterMeta": (".index_manager", "ChapterMeta"),
"SceneMeta": (".index_manager", "SceneMeta"),
"ReviewMetrics": (".index_manager", "ReviewMetrics"),
"RelationshipEventMeta": (".index_manager", "RelationshipEventMeta"),
# RAG Adapter
"RAGAdapter": (".rag_adapter", "RAGAdapter"),
"SearchResult": (".rag_adapter", "SearchResult"),
"ContextManager": (".context_manager", "ContextManager"),
"ContextRanker": (".context_ranker", "ContextRanker"),
"SnapshotManager": (".snapshot_manager", "SnapshotManager"),
"QueryRouter": (".query_router", "QueryRouter"),
# Style Sampler
"StyleSampler": (".style_sampler", "StyleSampler"),
"StyleSample": (".style_sampler", "StyleSample"),
"SceneType": (".style_sampler", "SceneType"),
}
def __getattr__(name: str) -> Any: # pragma: no cover
if name not in _LAZY_EXPORTS:
raise AttributeError(name)
module_path, attr = _LAZY_EXPORTS[name]
module = import_module(module_path, __name__)
value = getattr(module, attr)
globals()[name] = value # cache
return value
def __dir__() -> list[str]: # pragma: no cover
return sorted(set(list(globals().keys()) + list(_LAZY_EXPORTS.keys())))
+495
View File
@@ -0,0 +1,495 @@
#!/usr/bin/env python3
# -*- coding: utf-8 -*-
"""
Data Modules - API 客户端 (v5.4,v5.0 OpenAI 兼容接口沿用)
支持两种 API 类型:
1. openai: OpenAI 兼容的 /v1/embeddings 和 /v1/rerank 接口
- 适用于: OpenAI, Jina, Cohere, vLLM, Ollama 等
2. modal: Modal 自定义接口格式
- 适用于: 自部署的 Modal 服务
配置示例 (config.py):
embed_api_type = "openai"
embed_base_url = "https://api.openai.com/v1"
embed_model = "text-embedding-3-small"
embed_api_key = "sk-xxx"
rerank_api_type = "openai" # Jina/Cohere 也使用此类型
rerank_base_url = "https://api.jina.ai/v1"
rerank_model = "jina-reranker-v2-base-multilingual"
rerank_api_key = "jina_xxx"
"""
import asyncio
import aiohttp
import time
from typing import List, Dict, Any, Optional
from dataclasses import dataclass
from .config import get_config
@dataclass
class APIStats:
"""API 调用统计"""
total_calls: int = 0
total_time: float = 0.0
errors: int = 0
class EmbeddingAPIClient:
"""
通用 Embedding API 客户端
支持 OpenAI 兼容接口 (/v1/embeddings) 和 Modal 自定义接口
"""
def __init__(self, config=None):
self.config = config or get_config()
self.sem = asyncio.Semaphore(self.config.embed_concurrency)
self.stats = APIStats()
self._warmed_up = False
self._session: Optional[aiohttp.ClientSession] = None
self.last_error_status: Optional[int] = None
self.last_error_message: str = ""
async def _get_session(self) -> aiohttp.ClientSession:
if self._session is None or self._session.closed:
connector = aiohttp.TCPConnector(limit=200, limit_per_host=100)
self._session = aiohttp.ClientSession(connector=connector)
return self._session
async def close(self):
if self._session and not self._session.closed:
await self._session.close()
def _build_headers(self) -> Dict[str, str]:
"""构建请求头"""
headers = {"Content-Type": "application/json"}
if self.config.embed_api_key:
headers["Authorization"] = f"Bearer {self.config.embed_api_key}"
return headers
def _build_url(self) -> str:
"""构建请求 URL"""
base_url = self.config.embed_base_url.rstrip("/")
if self.config.embed_api_type == "openai":
# OpenAI 兼容: /v1/embeddings
if not base_url.endswith("/embeddings"):
if base_url.endswith("/v1"):
return f"{base_url}/embeddings"
return f"{base_url}/v1/embeddings"
return base_url
else:
# Modal 自定义接口: 直接使用配置的 URL
return base_url
def _build_payload(self, texts: List[str]) -> Dict[str, Any]:
"""构建请求体"""
if self.config.embed_api_type == "openai":
return {
"input": texts,
"model": self.config.embed_model,
"encoding_format": "float"
}
else:
# Modal 格式
return {
"input": texts,
"model": self.config.embed_model
}
def _parse_response(self, data: Dict[str, Any]) -> Optional[List[List[float]]]:
"""解析响应"""
if self.config.embed_api_type == "openai":
# OpenAI 格式: {"data": [{"embedding": [...], "index": 0}, ...]}
if "data" in data:
# 按 index 排序,确保顺序正确
sorted_data = sorted(data["data"], key=lambda x: x.get("index", 0))
return [item["embedding"] for item in sorted_data]
return None
else:
# Modal 格式: {"data": [{"embedding": [...]}, ...]}
if "data" in data:
return [item["embedding"] for item in data["data"]]
return None
async def embed(self, texts: List[str]) -> Optional[List[List[float]]]:
"""调用 Embedding 服务(带重试机制)"""
if not texts:
return []
timeout = self.config.cold_start_timeout if not self._warmed_up else self.config.normal_timeout
max_retries = getattr(self.config, 'api_max_retries', 3)
base_delay = getattr(self.config, 'api_retry_delay', 1.0)
async with self.sem:
start = time.time()
session = await self._get_session()
for attempt in range(max_retries):
try:
url = self._build_url()
headers = self._build_headers()
payload = self._build_payload(texts)
async with session.post(
url,
json=payload,
headers=headers,
timeout=aiohttp.ClientTimeout(total=timeout)
) as resp:
if resp.status == 200:
text = await resp.text()
import json as json_module
data = json_module.loads(text)
embeddings = self._parse_response(data)
if embeddings:
self.stats.total_calls += 1
self.stats.total_time += time.time() - start
self._warmed_up = True
self.last_error_status = None
self.last_error_message = ""
return embeddings
# 可重试的状态码: 429 (限流), 500, 502, 503, 504
if resp.status in (429, 500, 502, 503, 504) and attempt < max_retries - 1:
delay = base_delay * (2 ** attempt) # 指数退避
print(f"[WARN] Embed {resp.status}, retrying in {delay:.1f}s ({attempt + 1}/{max_retries})")
await asyncio.sleep(delay)
continue
self.stats.errors += 1
err_text = await resp.text()
self.last_error_status = int(resp.status)
self.last_error_message = str(err_text[:200])
print(f"[ERR] Embed {resp.status}: {err_text[:200]}")
return None
except asyncio.TimeoutError:
if attempt < max_retries - 1:
delay = base_delay * (2 ** attempt)
print(f"[WARN] Embed timeout, retrying in {delay:.1f}s ({attempt + 1}/{max_retries})")
await asyncio.sleep(delay)
continue
self.stats.errors += 1
self.last_error_status = None
self.last_error_message = f"Timeout after {max_retries} attempts"
print(f"[ERR] Embed: Timeout after {max_retries} attempts")
return None
except Exception as e:
if attempt < max_retries - 1:
delay = base_delay * (2 ** attempt)
print(f"[WARN] Embed error: {e}, retrying in {delay:.1f}s ({attempt + 1}/{max_retries})")
await asyncio.sleep(delay)
continue
self.stats.errors += 1
self.last_error_status = None
self.last_error_message = str(e)
print(f"[ERR] Embed: {e}")
return None
return None
async def embed_batch(
self, texts: List[str], *, skip_failures: bool = True
) -> List[Optional[List[float]]]:
"""
分批 Embedding
Args:
texts: 要嵌入的文本列表
skip_failures: True 时失败的文本返回 None;False 时任一失败则整体返回空列表
Returns:
与 texts 等长的列表,成功的位置是向量,失败的位置是 None
"""
if not texts:
return []
all_embeddings: List[Optional[List[float]]] = []
batch_size = self.config.embed_batch_size
batches = [texts[i:i + batch_size] for i in range(0, len(texts), batch_size)]
tasks = [self.embed(batch) for batch in batches]
results = await asyncio.gather(*tasks)
for batch_idx, result in enumerate(results):
actual_batch_size = len(batches[batch_idx])
if result and len(result) == actual_batch_size:
all_embeddings.extend(result)
else:
if not skip_failures:
print(f"[WARN] Embed batch {batch_idx} failed, aborting all")
return []
print(f"[WARN] Embed batch {batch_idx} failed, marking {actual_batch_size} items as None")
all_embeddings.extend([None] * actual_batch_size)
return all_embeddings[:len(texts)]
async def warmup(self):
"""预热服务"""
await self.embed(["test"])
self._warmed_up = True
class RerankAPIClient:
"""
通用 Rerank API 客户端
支持 OpenAI 兼容接口 (Jina/Cohere 格式) 和 Modal 自定义接口
"""
def __init__(self, config=None):
self.config = config or get_config()
self.sem = asyncio.Semaphore(self.config.rerank_concurrency)
self.stats = APIStats()
self._warmed_up = False
self._session: Optional[aiohttp.ClientSession] = None
async def _get_session(self) -> aiohttp.ClientSession:
if self._session is None or self._session.closed:
connector = aiohttp.TCPConnector(limit=200, limit_per_host=100)
self._session = aiohttp.ClientSession(connector=connector)
return self._session
async def close(self):
if self._session and not self._session.closed:
await self._session.close()
def _build_headers(self) -> Dict[str, str]:
"""构建请求头"""
headers = {"Content-Type": "application/json"}
if self.config.rerank_api_key:
headers["Authorization"] = f"Bearer {self.config.rerank_api_key}"
return headers
def _build_url(self) -> str:
"""构建请求 URL"""
base_url = self.config.rerank_base_url.rstrip("/")
if self.config.rerank_api_type == "openai":
# Jina/Cohere 兼容: /v1/rerank
if not base_url.endswith("/rerank"):
if base_url.endswith("/v1"):
return f"{base_url}/rerank"
return f"{base_url}/v1/rerank"
return base_url
else:
# Modal 自定义接口
return base_url
def _build_payload(self, query: str, documents: List[str], top_n: Optional[int]) -> Dict[str, Any]:
"""构建请求体"""
if self.config.rerank_api_type == "openai":
# Jina/Cohere 格式
payload: Dict[str, Any] = {
"query": query,
"documents": documents,
"model": self.config.rerank_model
}
if top_n:
payload["top_n"] = top_n
return payload
else:
# Modal 格式
payload = {"query": query, "documents": documents}
if top_n:
payload["top_n"] = top_n
return payload
def _parse_response(self, data: Dict[str, Any]) -> List[Dict[str, Any]]:
"""解析响应"""
if self.config.rerank_api_type == "openai":
# Jina/Cohere 格式: {"results": [{"index": 0, "relevance_score": 0.9}, ...]}
return data.get("results", [])
else:
# Modal 格式: {"results": [...]}
return data.get("results", [])
async def rerank(
self,
query: str,
documents: List[str],
top_n: Optional[int] = None
) -> Optional[List[Dict[str, Any]]]:
"""调用 Rerank 服务(带重试机制)"""
if not documents:
return []
timeout = self.config.cold_start_timeout if not self._warmed_up else self.config.normal_timeout
max_retries = getattr(self.config, 'api_max_retries', 3)
base_delay = getattr(self.config, 'api_retry_delay', 1.0)
async with self.sem:
start = time.time()
session = await self._get_session()
for attempt in range(max_retries):
try:
url = self._build_url()
headers = self._build_headers()
payload = self._build_payload(query, documents, top_n)
async with session.post(
url,
json=payload,
headers=headers,
timeout=aiohttp.ClientTimeout(total=timeout)
) as resp:
if resp.status == 200:
data = await resp.json()
self.stats.total_calls += 1
self.stats.total_time += time.time() - start
self._warmed_up = True
return self._parse_response(data)
# 可重试的状态码
if resp.status in (429, 500, 502, 503, 504) and attempt < max_retries - 1:
delay = base_delay * (2 ** attempt)
print(f"[WARN] Rerank {resp.status}, retrying in {delay:.1f}s ({attempt + 1}/{max_retries})")
await asyncio.sleep(delay)
continue
self.stats.errors += 1
err_text = await resp.text()
print(f"[ERR] Rerank {resp.status}: {err_text[:200]}")
return None
except asyncio.TimeoutError:
if attempt < max_retries - 1:
delay = base_delay * (2 ** attempt)
print(f"[WARN] Rerank timeout, retrying in {delay:.1f}s ({attempt + 1}/{max_retries})")
await asyncio.sleep(delay)
continue
self.stats.errors += 1
print(f"[ERR] Rerank: Timeout after {max_retries} attempts")
return None
except Exception as e:
if attempt < max_retries - 1:
delay = base_delay * (2 ** attempt)
print(f"[WARN] Rerank error: {e}, retrying in {delay:.1f}s ({attempt + 1}/{max_retries})")
await asyncio.sleep(delay)
continue
self.stats.errors += 1
print(f"[ERR] Rerank: {e}")
return None
return None
async def warmup(self):
"""预热服务"""
await self.rerank("test", ["doc1", "doc2"])
self._warmed_up = True
class ModalAPIClient:
"""
统一 API 客户端 (兼容旧接口)
整合 Embedding + Rerank 客户端,保持向后兼容
"""
def __init__(self, config=None):
self.config = config or get_config()
self._embed_client = EmbeddingAPIClient(self.config)
self._rerank_client = RerankAPIClient(self.config)
# 兼容旧代码的信号量
self.sem_embed = self._embed_client.sem
self.sem_rerank = self._rerank_client.sem
self._warmed_up = {"embed": False, "rerank": False}
self._session: Optional[aiohttp.ClientSession] = None
@property
def stats(self) -> Dict[str, APIStats]:
return {
"embed": self._embed_client.stats,
"rerank": self._rerank_client.stats
}
async def _get_session(self) -> aiohttp.ClientSession:
# 复用 embed client 的 session
return await self._embed_client._get_session()
async def close(self):
await self._embed_client.close()
await self._rerank_client.close()
# ==================== 预热 ====================
async def warmup(self):
"""预热 Embedding 和 Rerank 服务"""
print("[WARMUP] Warming up Embed + Rerank...")
start = time.time()
tasks = [self._warmup_embed(), self._warmup_rerank()]
results = await asyncio.gather(*tasks, return_exceptions=True)
for name, result in zip(["Embed", "Rerank"], results):
if isinstance(result, Exception):
print(f" [FAIL] {name}: {result}")
else:
print(f" [OK] {name} ready")
print(f"[WARMUP] Done in {time.time() - start:.1f}s")
async def _warmup_embed(self):
await self._embed_client.warmup()
self._warmed_up["embed"] = True
async def _warmup_rerank(self):
await self._rerank_client.warmup()
self._warmed_up["rerank"] = True
# ==================== Embedding API ====================
async def embed(self, texts: List[str]) -> Optional[List[List[float]]]:
"""调用 Embedding 服务"""
return await self._embed_client.embed(texts)
async def embed_batch(
self, texts: List[str], *, skip_failures: bool = True
) -> List[Optional[List[float]]]:
"""分批 Embedding"""
return await self._embed_client.embed_batch(texts, skip_failures=skip_failures)
# ==================== Rerank API ====================
async def rerank(
self,
query: str,
documents: List[str],
top_n: Optional[int] = None
) -> Optional[List[Dict[str, Any]]]:
"""调用 Rerank 服务"""
return await self._rerank_client.rerank(query, documents, top_n)
# ==================== 统计 ====================
def print_stats(self):
print("\n[API STATS]")
for name, stats in self.stats.items():
if stats.total_calls > 0:
avg_time = stats.total_time / stats.total_calls
print(f" {name.upper()}: {stats.total_calls} calls, "
f"{stats.total_time:.1f}s total, "
f"{avg_time:.2f}s avg, "
f"{stats.errors} errors")
# 全局客户端
_client: Optional[ModalAPIClient] = None
def get_client(config=None) -> ModalAPIClient:
global _client
if _client is None or config is not None:
_client = ModalAPIClient(config)
return _client
@@ -0,0 +1,938 @@
#!/usr/bin/env python3
# -*- coding: utf-8 -*-
"""
Chapter Analyzer - 章节自动解构与模式提取
功能:
1. 章节内容解析(场景分割、人物提取)
2. 张力曲线分析(情绪压强跟踪)
3. 爽点模式识别(打脸、降维、闭环、僭越)
4. 套路结构提取(铺垫→压抑→爆发→反转)
5. 学习成果输出(可复用模式)
用法:
python chapter_analyzer.py --chapter-file "正文/第0003章-当众打脸.md" --project-root . --learn
"""
import re
import json
import hashlib
import sqlite3
from pathlib import Path
from dataclasses import dataclass, field, asdict
from typing import List, Dict, Any, Optional, Tuple
from datetime import datetime
from enum import Enum
import asyncio
try:
from runtime_compat import enable_windows_utf8_stdio
except ImportError:
enable_windows_utf8_stdio = lambda: None
class PatternType(Enum):
"""模式类型"""
HOOK = "hook" # 钩子
PACING = "pacing" # 节奏
DIALOGUE = "dialogue" # 对话
PAYOFF = "payoff" # 兑现
EMOTION = "emotion" # 情绪
TENSION_BUILD = "tension_build" # 压强积累
RELEASE = "release" # 释放
TWIST = "twist" # 反转
class CatharsisModel(Enum):
"""爽感模型"""
TABOO_TRANSGRESSION = "taboo_transgression" # 禁忌僭越
OVERKILL_REVERSAL = "overkill_reversal" # 降维打击
COGNITIVE_CLOSURE = "cognitive_closure" # 认知闭环
@dataclass
class TensionPoint:
"""张力点"""
position: int # 位置(字符偏移)
tension: float # 张力值 0-1
event_type: str # 事件类型
description: str # 描述
@dataclass
class HotSpot:
"""爽点"""
position: int
span: int # 持续长度
catharsis_type: str # 爽感类型
intensity: float # 强度 0-1
description: str
@dataclass
class ChapterStructure:
"""章节结构"""
setup: int = 0 # 铺垫段落长度
suppress: int = 0 # 压抑段落长度
release: int = 0 # 释放段落长度
twist: int = 0 # 反转段落长度
@dataclass
class LearnedPattern:
"""学习到的模式"""
pattern_id: str
pattern_type: str
title: str
description: str
tension_curve: List[List[float]] # [[position, tension], ...]
catharsis_model: str
structure: Dict[str, int]
hot_spots: List[List[Any]]
style_tags: List[str]
source_project: str
source_chapter: int
learned_at: str
usage_count: int = 0
metadata: Dict[str, Any] = field(default_factory=dict)
@dataclass
class ChapterAnalysisResult:
"""章节分析结果"""
chapter: int
title: str
word_count: int
scene_count: int
character_count: int
# 张力分析
tension_curve: List[List[float]] # [[position, tension], ...]
tension_avg: float
tension_peak: float
tension_peak_position: int
# 爽点分析
hot_spots: List[HotSpot]
catharsis_model: str
# 结构分析
structure: ChapterStructure
# 模式
detected_patterns: List[str]
# 风格标签
style_tags: List[str]
# 人物出场
characters: List[str]
# 建议
suggestions: List[str]
class ChapterAnalyzer:
"""
章节分析器
自动解构章节,提取:
- 张力曲线
- 爽点位置和类型
- 套路结构
- 可复用模式
"""
# 爽点关键词
HOT_SPOT_KEYWORDS = {
CatharsisModel.OVERKILL_REVERSAL: [
"碾压", "秒杀", "不堪一击", "一招", "跪下", "颤抖", "惊恐",
"脸色大变", "难以置信", "怎么可能", "废物", "蝼蚁", "打脸",
"当众", "所有人", "目瞪口呆", "鸦雀无声", "死寂"
],
CatharsisModel.TABOO_TRANSGRESSION: [
"禁忌", "僭越", "突破底线", "危险", "禁忌", "邪魅", "诱惑",
"堕落", "黑化", "疯狂", "失控", "暴走", "越界", "禁忌"
],
CatharsisModel.COGNITIVE_CLOSURE: [
"原来", "竟然", "真相", "伏笔", "恍然大悟", "所有一切",
"早该", "早就在", "铺垫", "埋下", "算计", "布局"
]
}
# 情绪压强关键词
TENSION_KEYWORDS = {
# 压强上升
"rise": ["紧张", "危机", "危险", "困境", "难题", "冲突", "对峙",
"杀意", "阴谋", "威胁", "危机", "悬念", "未知"],
# 压强高峰
"peak": ["爆发", "突破", "反击", "反转", "真相", "高潮", "决战",
"绝杀", "逆转", "翻盘", "打脸", "碾压"],
# 压强下降
"fall": ["松了口气", "终于", "安心", "胜利", "结束", "平息", "安然"]
}
def __init__(self, project_root: Optional[Path] = None):
self.project_root = Path(project_root) if project_root else Path.cwd()
def analyze_chapter(
self,
chapter_file: Path,
chapter_num: Optional[int] = None
) -> ChapterAnalysisResult:
"""
分析章节
Args:
chapter_file: 章节文件路径
chapter_num: 章节号(从文件名推断或手动指定)
Returns:
章节分析结果
"""
content = chapter_file.read_text(encoding="utf-8")
# 解析章节号
if chapter_num is None:
chapter_num = self._extract_chapter_num(chapter_file.name)
# 解析标题
title = self._extract_title(chapter_file.name, content)
# 统计字数
word_count = len(content)
# 场景分割
scenes = self._split_scenes(content)
scene_count = len(scenes)
# 人物提取
characters = self._extract_characters(content, scenes)
# 张力曲线分析
tension_curve = self._analyze_tension_curve(content, scenes)
tension_avg = sum(p[1] for p in tension_curve) / len(tension_curve) if tension_curve else 0
peak = max(tension_curve, key=lambda x: x[1]) if tension_curve else (0, 0)
tension_peak = peak[1]
tension_peak_position = peak[0]
# 爽点分析
hot_spots = self._detect_hot_spots(content, scenes)
catharsis_model = self._detect_catharsis_model(hot_spots)
# 结构分析
structure = self._analyze_structure(content, scenes, hot_spots)
# 模式检测
detected_patterns = self._detect_patterns(
content, scenes, tension_curve, hot_spots, structure
)
# 风格标签
style_tags = self._extract_style_tags(content, tension_curve, hot_spots)
# 建议
suggestions = self._generate_suggestions(
tension_curve, hot_spots, structure, detected_patterns
)
return ChapterAnalysisResult(
chapter=chapter_num,
title=title,
word_count=word_count,
scene_count=scene_count,
character_count=len(characters),
tension_curve=tension_curve,
tension_avg=tension_avg,
tension_peak=tension_peak,
tension_peak_position=tension_peak_position,
hot_spots=hot_spots,
catharsis_model=catharsis_model,
structure=structure,
detected_patterns=detected_patterns,
style_tags=style_tags,
characters=characters,
suggestions=suggestions
)
def _extract_chapter_num(self, filename: str) -> int:
"""从文件名提取章节号"""
match = re.search(r'第(\d+)[章节]', filename)
if match:
return int(match.group(1))
return 0
def _extract_title(self, filename: str, content: str) -> str:
"""提取标题"""
# 尝试从文件名提取
match = re.search(r'第\d+章[章节]-?(.+)', filename)
if match:
return match.group(1).replace('.md', '').strip()
# 尝试从内容第一行提取
lines = content.split('\n')
for line in lines[:5]:
line = line.strip()
if line.startswith('#'):
return line.lstrip('#').strip()
return filename.replace('.md', '')
def _split_scenes(self, content: str) -> List[Dict[str, Any]]:
"""
分割场景
Returns:
List[{"start": int, "end": int, "location": str, "characters": []}]
"""
scenes = []
lines = content.split('\n')
current_scene = {
"start": 0,
"end": 0,
"location": "未知",
"characters": [],
"lines": []
}
location_patterns = [
r'^---+$', # 分割线
r'^【(.+?)】', # 【场景名】
r'^((.+?))', # (场景名)
]
for i, line in enumerate(lines):
# 检测场景分隔
is_divider = any(re.match(p, line.strip()) for p in location_patterns)
if is_divider and current_scene["lines"]:
# 保存当前场景
current_scene["end"] = sum(len(l) + 1 for l in current_scene["lines"][:-1])
scenes.append(current_scene)
# 新场景
location_match = re.search(r'【(.+?)】|((.+?))', line)
current_scene = {
"start": sum(len(l) + 1 for l in lines[:i]) + 1,
"end": 0,
"location": location_match.group(1) if location_match else "未知",
"characters": [],
"lines": []
}
current_scene["lines"].append(line)
# 保存最后一个场景
if current_scene["lines"]:
current_scene["end"] = len(content)
scenes.append(current_scene)
return scenes if scenes else [{"start": 0, "end": len(content), "location": "未知", "characters": [], "lines": content.split('\n')}]
def _extract_characters(
self,
content: str,
scenes: List[Dict[str, Any]]
) -> List[str]:
"""提取人物列表"""
# 简单实现:提取引号内的对话人
characters = set()
# 匹配 "XXX说" 格式
dialogue_pattern = re.compile(r'^"?([^"说]{2,5})"?[说问道喊叫笑骂冷哼]')
for scene in scenes:
for line in scene["lines"]:
match = dialogue_pattern.match(line.strip())
if match:
name = match.group(1).strip()
if name and len(name) <= 5:
characters.add(name)
return list(characters)
def _analyze_tension_curve(
self,
content: str,
scenes: List[Dict[str, Any]]
) -> List[List[float]]:
"""
分析张力曲线
Returns:
[[position, tension], ...] - 位置(字数)和张力值(0-1)
"""
curve = []
total_len = len(content)
segment_size = 200 # 每200字一个采样点
content_lower = content # 网文不区分大小写
for pos in range(0, total_len, segment_size):
segment = content_lower[pos:pos + segment_size]
tension = 0.0
# 检查压强上升关键词
for kw in self.TENSION_KEYWORDS["rise"]:
if kw in segment:
tension += 0.2
# 检查压强高峰关键词
for kw in self.TENSION_KEYWORDS["peak"]:
if kw in segment:
tension += 0.4
# 检查压强下降关键词
for kw in self.TENSION_KEYWORDS["fall"]:
if kw in segment:
tension -= 0.15
# 限制范围
tension = max(0.0, min(1.0, tension))
curve.append([pos, tension])
# 平滑曲线
curve = self._smooth_curve(curve)
return curve
def _smooth_curve(self, curve: List[List[float]]) -> List[List[float]]:
"""平滑张力曲线"""
if len(curve) < 3:
return curve
smoothed = []
for i, point in enumerate(curve):
if i == 0 or i == len(curve) - 1:
smoothed.append(point)
else:
# 移动平均
avg_pos = (curve[i-1][0] + point[0] + curve[i+1][0]) / 3
avg_tension = (curve[i-1][1] + point[1] + curve[i+1][1]) / 3
smoothed.append([avg_pos, avg_tension])
return smoothed
def _detect_hot_spots(
self,
content: str,
scenes: List[Dict[str, Any]]
) -> List[HotSpot]:
"""检测爽点"""
hot_spots = []
for model_type, keywords in self.HOT_SPOT_KEYWORDS.items():
for keyword in keywords:
# 查找所有出现位置
start = 0
while True:
pos = content.find(keyword, start)
if pos == -1:
break
# 计算附近区域的强度
context_start = max(0, pos - 100)
context_end = min(len(content), pos + 100)
context = content[context_start:context_end]
# 统计上下文中的情绪词数量
intensity = 0.5 # 基础强度
for kw_set in self.HOT_SPOT_KEYWORDS.values():
for kw in kw_set:
if kw in context:
intensity += 0.1
intensity = min(1.0, intensity)
# 避免重叠
span = len(keyword)
is_overlap = any(
abs(pos - hs.position) < span
for hs in hot_spots
)
if not is_overlap:
hot_spots.append(HotSpot(
position=pos,
span=span,
catharsis_type=model_type.value,
intensity=intensity,
description=f"发现「{keyword}」"
))
start = pos + 1
# 按位置排序
hot_spots.sort(key=lambda x: x.position)
return hot_spots
def _detect_catharsis_model(
self,
hot_spots: List[HotSpot]
) -> str:
"""检测主导爽感模型"""
if not hot_spots:
return "unknown"
model_counts = {}
model_intensity = {}
for hs in hot_spots:
model = hs.catharsis_type
model_counts[model] = model_counts.get(model, 0) + 1
model_intensity[model] = model_intensity.get(model, 0) + hs.intensity
# 综合评分:出现次数 * 0.4 + 强度 * 0.6
scores = {
m: model_counts[m] * 0.4 + model_intensity[m] * 0.6
for m in model_counts
}
return max(scores, key=scores.get) if scores else "unknown"
def _analyze_structure(
self,
content: str,
scenes: List[Dict[str, Any]],
hot_spots: List[HotSpot]
) -> ChapterStructure:
"""分析章节结构"""
total_len = len(content)
if not hot_spots:
# 无爽点:均匀分布
segment = total_len // 4
return ChapterStructure(
setup=segment,
suppress=segment,
release=segment,
twist=segment
)
# 找第一个爽点位置作为分界
first_hot = hot_spots[0].position
last_hot = hot_spots[-1].position
# 铺垫:第一个爽点之前
setup = first_hot
# 压抑+释放:根据爽点密集程度判断
# 简化:中间区域前40%压抑,后40%释放
middle_start = first_hot
middle_end = last_hot + 100
middle_len = middle_end - middle_start
suppress = int(middle_len * 0.4)
release = int(middle_len * 0.4)
twist = total_len - middle_end
return ChapterStructure(
setup=max(0, setup),
suppress=max(0, suppress),
release=max(0, release),
twist=max(0, twist)
)
def _detect_patterns(
self,
content: str,
scenes: List[Dict[str, Any]],
tension_curve: List[List[float]],
hot_spots: List[HotSpot],
structure: ChapterStructure
) -> List[str]:
"""检测章节中的模式"""
patterns = []
# 检测钩子模式
if tension_curve and tension_curve[0][1] > 0.3:
patterns.append("开篇高能钩")
elif "?" in content[:200] or "!" in content[:200]:
patterns.append("悬念钩子")
# 检测打脸模式
if any(hs.catharsis_type == "overkill_reversal" for hs in hot_spots):
patterns.append("打脸反转")
# 检测升级模式
if any(kw in content for kw in ["突破", "晋升", "升级", "进阶"]):
patterns.append("境界突破")
# 检测装逼模式
if any(kw in content for kw in ["冷笑", "不屑", "蝼蚁", "可笑"]):
patterns.append("装逼打脸")
# 张力曲线形状检测
if tension_curve:
# 持续上升 = 压抑型
tensions = [p[1] for p in tension_curve]
if tensions == sorted(tensions) and tensions[-1] - tensions[0] > 0.5:
patterns.append("单线压抑")
# 波动大 = 节奏快
elif max(tensions) - min(tensions) > 0.6:
patterns.append("爽点密集")
return patterns
def _extract_style_tags(
self,
content: str,
tension_curve: List[List[float]],
hot_spots: List[HotSpot]
) -> List[str]:
"""提取风格标签"""
tags = []
# 基于字数
word_count = len(content)
if word_count < 2000:
tags.append("短小精悍")
elif word_count > 5000:
tags.append("长篇巨制")
# 基于爽点密度
if tension_curve:
avg_tension = sum(p[1] for p in tension_curve) / len(tension_curve)
if avg_tension > 0.5:
tags.append("情绪压强高")
elif avg_tension < 0.2:
tags.append("节奏舒缓")
# 基于爽点类型
if hot_spots:
model = self._detect_catharsis_model(hot_spots)
if model == "overkill_reversal":
tags.append("爽点密集")
elif model == "taboo_transgression":
tags.append("边缘拉扯")
elif model == "cognitive_closure":
tags.append("多线收束")
# 基于对话比例
dialogue_count = content.count('"') + content.count('"')
dialogue_ratio = dialogue_count / max(word_count, 1)
if dialogue_ratio > 0.3:
tags.append("对话驱动")
elif dialogue_ratio < 0.1:
tags.append("叙事为主")
return tags
def _generate_suggestions(
self,
tension_curve: List[List[float]],
hot_spots: List[HotSpot],
structure: ChapterStructure,
patterns: List[str]
) -> List[str]:
"""生成改进建议"""
suggestions = []
# 张力曲线建议
if tension_curve:
avg = sum(p[1] for p in tension_curve) / len(tension_curve)
if avg < 0.2:
suggestions.append("张力整体偏低,建议增加危机感或悬念")
elif avg > 0.7:
suggestions.append("张力持续偏高,建议适当释放避免读者疲劳")
# 爽点建议
if not hot_spots:
suggestions.append("未检测到明显爽点,建议增加打脸/反转情节")
# 结构建议
if structure.setup < 500:
suggestions.append("铺垫不足,建议增加背景/动机描写")
# 节奏建议
if "单线压抑" in patterns:
suggestions.append("压抑较长,建议穿插小高潮保持节奏")
return suggestions
def learn_pattern(
self,
analysis: ChapterAnalysisResult,
layer: str = "project"
) -> LearnedPattern:
"""
从分析结果提取可学习模式
Args:
analysis: 章节分析结果
layer: 存储层级 ("project" 或 "system")
Returns:
学习到的模式
"""
pattern_id = hashlib.md5(
f"{self.project_root.name}_{analysis.chapter}_{datetime.now().isoformat()}".encode()
).hexdigest()[:16]
return LearnedPattern(
pattern_id=pattern_id,
pattern_type=self._infer_pattern_type(analysis),
title=f"第{analysis.chapter}章模式: {analysis.title}",
description=self._generate_description(analysis),
tension_curve=analysis.tension_curve,
catharsis_model=analysis.catharsis_model,
structure=asdict(analysis.structure),
hot_spots=[[hs.position, hs.description] for hs in analysis.hot_spots],
style_tags=analysis.style_tags,
source_project=self.project_root.name,
source_chapter=analysis.chapter,
learned_at=datetime.now().isoformat(),
metadata={
"word_count": analysis.word_count,
"scene_count": analysis.scene_count,
"character_count": analysis.character_count,
"detected_patterns": analysis.detected_patterns,
"suggestions": analysis.suggestions
}
)
def _infer_pattern_type(self, analysis: ChapterAnalysisResult) -> str:
"""推断模式类型"""
if "打脸反转" in analysis.detected_patterns:
return PatternType.HOOK.value
elif "境界突破" in analysis.detected_patterns:
return PatternType.PAYOFF.value
elif analysis.tension_avg > 0.5:
return PatternType.TENSION_BUILD.value
else:
return PatternType.PACING.value
def _generate_description(self, analysis: ChapterAnalysisResult) -> str:
"""生成模式描述"""
parts = []
if analysis.detected_patterns:
parts.append(f"包含模式: {', '.join(analysis.detected_patterns[:3])}")
if analysis.style_tags:
parts.append(f"风格标签: {', '.join(analysis.style_tags[:3])}")
parts.append(f"字数: {analysis.word_count},场景: {analysis.scene_count}")
if analysis.hot_spots:
parts.append(f"检测到 {len(analysis.hot_spots)} 个爽点")
return "; ".join(parts)
def save_learned_pattern(
self,
pattern: LearnedPattern,
layer: str = "project"
) -> bool:
"""保存学习到的模式"""
try:
if layer == "project":
return self._save_to_project_db(pattern)
else:
return self._save_to_system(pattern)
except Exception as e:
print(f"Error saving pattern: {e}")
return False
def _init_project_learned_db(self) -> Path:
"""初始化项目学习库"""
db_dir = self.project_root / ".noma" / "rag"
db_dir.mkdir(parents=True, exist_ok=True)
db_path = db_dir / "learned.db"
conn = sqlite3.connect(str(db_path))
cursor = conn.cursor()
cursor.execute("""
CREATE TABLE IF NOT EXISTS learned_patterns (
pattern_id TEXT PRIMARY KEY,
pattern_type TEXT NOT NULL,
title TEXT NOT NULL,
description TEXT,
tension_curve TEXT,
catharsis_model TEXT,
structure TEXT,
hot_spots TEXT,
style_tags TEXT,
source_project TEXT,
source_chapter INTEGER,
learned_at TEXT,
usage_count INTEGER DEFAULT 0,
metadata TEXT
)
""")
conn.commit()
conn.close()
return db_path
def _save_to_project_db(self, pattern: LearnedPattern) -> bool:
"""保存到项目数据库"""
db_path = self._init_project_learned_db()
conn = sqlite3.connect(str(db_path))
cursor = conn.cursor()
cursor.execute("""
INSERT OR REPLACE INTO learned_patterns
(pattern_id, pattern_type, title, description, tension_curve,
catharsis_model, structure, hot_spots, style_tags,
source_project, source_chapter, learned_at, usage_count, metadata)
VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)
""", (
pattern.pattern_id,
pattern.pattern_type,
pattern.title,
pattern.description,
json.dumps(pattern.tension_curve),
pattern.catharsis_model,
json.dumps(pattern.structure),
json.dumps(pattern.hot_spots),
json.dumps(pattern.style_tags),
pattern.source_project,
pattern.source_chapter,
pattern.learned_at,
pattern.usage_count,
json.dumps(pattern.metadata)
))
conn.commit()
conn.close()
return True
def _save_to_system(self, pattern: LearnedPattern) -> bool:
"""保存到系统共享库"""
# 需要系统根目录
system_root = self._resolve_system_root()
learned_dir = system_root / "rag" / "learned"
learned_dir.mkdir(parents=True, exist_ok=True)
pattern_file = learned_dir / f"{pattern.pattern_id}.json"
pattern_data = asdict(pattern)
pattern_file.write_text(
json.dumps(pattern_data, ensure_ascii=False, indent=2),
encoding="utf-8"
)
return True
def _resolve_system_root(self) -> Path:
"""解析系统根目录"""
# 与 project_root 同级的 .noma
sibling = self.project_root.parent / ".noma"
if sibling.exists():
return sibling
return self.project_root / ".noma"
# ==================== CLI 接口 ====================
def main():
import argparse
import sys
if sys.platform == "win32":
enable_windows_utf8_stdio()
parser = argparse.ArgumentParser(
description="Chapter Analyzer - 章节自动解构与模式提取"
)
parser.add_argument("--project-root", type=str, default=".",
help="项目根目录")
parser.add_argument("--chapter-file", type=str, required=True,
help="章节文件路径")
parser.add_argument("--chapter-num", type=int,
help="章节号(可选,从文件名推断)")
parser.add_argument("--learn", action="store_true",
help="学习并存储模式")
parser.add_argument("--learn-layer", choices=["project", "system"],
default="project",
help="学习成果存储层级")
parser.add_argument("--output", type=str,
help="输出文件(JSON格式)")
args = parser.parse_args()
project_root = Path(args.project_root).resolve()
chapter_file = project_root / args.chapter_file
if not chapter_file.exists():
print(f"Error: Chapter file not found: {chapter_file}")
sys.exit(1)
# 分析章节
analyzer = ChapterAnalyzer(project_root)
result = analyzer.analyze_chapter(chapter_file, args.chapter_num)
# 输出结果
if args.output:
output_path = Path(args.output)
output_path.parent.mkdir(parents=True, exist_ok=True)
output_path.write_text(
json.dumps(asdict(result), ensure_ascii=False, indent=2),
encoding="utf-8"
)
print(f"Analysis saved to: {output_path}")
else:
# 打印摘要
print(f"\n{'='*60}")
print(f"章节分析报告: 第{result.chapter}章 - {result.title}")
print(f"{'='*60}")
print(f"\n基本信息:")
print(f" 字数: {result.word_count}")
print(f" 场景数: {result.scene_count}")
print(f" 人物数: {result.character_count}")
print(f" 人物列表: {', '.join(result.characters[:5]) or '无'}")
print(f"\n张力分析:")
print(f" 平均张力: {result.tension_avg:.2f}")
print(f" 峰值张力: {result.tension_peak:.2f} (位置: {result.tension_peak_position})")
print(f"\n爽点分析:")
print(f" 爽点数量: {len(result.hot_spots)}")
print(f" 主导模型: {result.catharsis_model}")
if result.hot_spots:
print(f" 主要爽点:")
for hs in result.hot_spots[:3]:
print(f" - [{hs.position}] {hs.description} ({hs.intensity:.2f})")
print(f"\n结构分析:")
print(f" 铺垫: {result.structure.setup}字")
print(f" 压抑: {result.structure.suppress}字")
print(f" 释放: {result.structure.release}字")
print(f" 反转: {result.structure.twist}字")
print(f"\n检测到的模式:")
for p in result.detected_patterns:
print(f" - {p}")
print(f"\n风格标签:")
for t in result.style_tags:
print(f" - {t}")
if result.suggestions:
print(f"\n改进建议:")
for s in result.suggestions:
print(f" - {s}")
# 学习模式
if args.learn:
pattern = analyzer.learn_pattern(result, args.learn_layer)
success = analyzer.save_learned_pattern(pattern, args.learn_layer)
if success:
print(f"\n✓ 模式已学习并存储到 {args.learn_layer} 层")
print(f" Pattern ID: {pattern.pattern_id}")
else:
print(f"\n✗ 模式存储失败")
sys.exit(1)
if __name__ == "__main__":
main()
+96
View File
@@ -0,0 +1,96 @@
#!/usr/bin/env python3
# -*- coding: utf-8 -*-
"""
CLI 参数兼容工具。
背景:
- data_modules 下的 CLI 普遍使用 argparse + subparsers。
- argparse 的全局参数(例如 --project-root)要求出现在子命令之前:
python -m data_modules.index_manager --project-root X get-core-entities
但实际写作流程里(skills/agents 文档、工具调用)经常把 --project-root 放在子命令之后:
python -m data_modules.index_manager get-core-entities --project-root X
这会直接报 "unrecognized arguments"(见 issues7 日志)。
这里提供一个轻量的 argv 预处理:把 --project-root 从任意位置提取出来并前置,
让原有 argparse 定义无需大改即可兼容两种写法。
"""
from __future__ import annotations
import json
import sys
from pathlib import Path
from typing import Any
from typing import List, Optional, Tuple
def _extract_flag_value(argv: List[str], flag: str) -> Tuple[Optional[str], List[str]]:
"""
Extract a flag value from argv.
Supports:
- --flag VALUE
- --flag=VALUE
Returns:
- (value, remaining_argv)
- value uses the *last* occurrence when repeated.
- if a dangling `--flag` has no value, it is kept in remaining_argv for argparse to raise.
"""
value: Optional[str] = None
rest: List[str] = []
i = 0
while i < len(argv):
token = argv[i]
if token == flag:
if i + 1 < len(argv):
value = argv[i + 1]
i += 2
continue
# Dangling flag; keep it so argparse can error out properly.
rest.append(token)
i += 1
continue
if token.startswith(flag + "="):
value = token.split("=", 1)[1]
i += 1
continue
rest.append(token)
i += 1
return value, rest
def normalize_global_project_root(argv: List[str], *, flag: str = "--project-root") -> List[str]:
"""
Normalize argv so a global `--project-root` (when present) is moved before subcommands.
This makes argparse+subparsers accept both:
- `... --project-root X cmd ...`
- `... cmd ... --project-root X`
"""
value, rest = _extract_flag_value(argv, flag)
if value is None:
return argv
return [flag, value] + rest
def load_json_arg(raw: str) -> Any:
"""
解析 CLI 传入的 JSON 参数,支持两种形式:
- 直接 JSON 字符串:'{"a":1}'
- @ 文件路径:'@data.json'(从文件读取 JSON,避免 shell 引号地狱)
- 特例:'@-' 表示从 stdin 读取
"""
if raw is None:
raise ValueError("missing json arg")
text = str(raw).strip()
if text.startswith("@"):
target = text[1:].strip()
if not target:
raise ValueError("invalid json arg: '@' without path")
if target == "-":
content = sys.stdin.read()
else:
content = Path(target).read_text(encoding="utf-8")
return json.loads(content)
return json.loads(text)
+69
View File
@@ -0,0 +1,69 @@
#!/usr/bin/env python3
# -*- coding: utf-8 -*-
"""
CLI output helpers for data_modules.
All CLI tools should emit JSON payloads via these helpers.
"""
from __future__ import annotations
import json
from dataclasses import dataclass
from typing import Any, Dict, Optional
@dataclass
class ErrorPayload:
code: str
message: str
suggestion: Optional[str] = None
details: Optional[Dict[str, Any]] = None
def build_success(data: Any = None, message: str = "ok", warnings: Optional[list] = None) -> Dict[str, Any]:
payload: Dict[str, Any] = {
"status": "success",
"message": message,
}
if data is not None:
payload["data"] = data
if warnings:
payload["warnings"] = warnings
return payload
def build_error(
code: str,
message: str,
suggestion: Optional[str] = None,
details: Optional[Dict[str, Any]] = None,
) -> Dict[str, Any]:
error: Dict[str, Any] = {
"code": code,
"message": message,
}
if suggestion:
error["suggestion"] = suggestion
if details:
error["details"] = details
return {
"status": "error",
"error": error,
}
def print_json(payload: Dict[str, Any]) -> None:
print(json.dumps(payload, ensure_ascii=False))
def print_success(data: Any = None, message: str = "ok", warnings: Optional[list] = None) -> None:
print_json(build_success(data=data, message=message, warnings=warnings))
def print_error(
code: str,
message: str,
suggestion: Optional[str] = None,
details: Optional[Dict[str, Any]] = None,
) -> None:
print_json(build_error(code=code, message=message, suggestion=suggestion, details=details))
+361
View File
@@ -0,0 +1,361 @@
#!/usr/bin/env python3
# -*- coding: utf-8 -*-
"""
Data Modules - 配置文件
API 配置通过环境变量读取(支持 .env 文件):
- EMBED_BASE_URL, EMBED_MODEL, EMBED_API_KEY
- RERANK_BASE_URL, RERANK_MODEL, RERANK_API_KEY
"""
import os
from pathlib import Path
from dataclasses import dataclass, field
from typing import Optional
from runtime_compat import normalize_windows_path
from .context_weights import TEMPLATE_WEIGHTS_DYNAMIC_DEFAULT
def _get_user_claude_root() -> Path:
raw = os.environ.get("NOMA_CLAUDE_HOME") or os.environ.get("CLAUDE_HOME")
if raw:
try:
return normalize_windows_path(raw).expanduser().resolve()
except Exception:
return normalize_windows_path(raw).expanduser()
return (Path.home() / ".claude").resolve()
def _load_dotenv_file(env_path: Path, *, override: bool = False) -> bool:
if not env_path.exists():
return False
try:
with open(env_path, "r", encoding="utf-8") as f:
for line in f:
line = line.strip()
if line and not line.startswith("#") and "=" in line:
key, _, value = line.partition("=")
key = key.strip()
value = value.strip()
if not key:
continue
# 默认不覆盖已有环境变量(保持“显式 > .env”优先级)
if override or key not in os.environ:
os.environ[key] = value
return True
except Exception:
return False
def _load_dotenv():
"""
加载 .env 文件(best-effort)。
约定:
- 项目级 `.env`(当前工作目录下)优先;
- 全局 `.env` 作为兜底:`~/.claude/novelmaster/.env`
"""
# 1) 当前目录(常见:用户从项目根目录执行)
_load_dotenv_file(Path.cwd() / ".env", override=False)
# 2) 用户级全局(常见:skills/agents 全局安装,API key 放这里最省心)
global_env = _get_user_claude_root() / "novelmaster" / ".env"
_load_dotenv_file(global_env, override=False)
def _load_project_dotenv(project_root: Path) -> None:
"""
加载某个项目根目录下的 `.env`(best-effort)。
优先加载顺序:.noma/config.env > .env
注意:不覆盖已存在环境变量,避免意外串台。
"""
project_root = Path(project_root)
# 优先加载 .noma/config.env(新版)
try:
_load_dotenv_file(project_root / ".noma" / "config.env", override=False)
except Exception:
pass
# 兼容旧版 .env
try:
_load_dotenv_file(project_root / ".env", override=False)
except Exception:
pass
_load_dotenv()
def _default_context_template_weights_dynamic() -> dict[str, dict[str, dict[str, float]]]:
return {
stage: {
template: dict(weights)
for template, weights in templates.items()
}
for stage, templates in TEMPLATE_WEIGHTS_DYNAMIC_DEFAULT.items()
}
@dataclass
class DataModulesConfig:
"""数据模块配置"""
# ================= 项目路径 =================
project_root: Path = field(default_factory=lambda: Path.cwd())
@property
def noma_dir(self) -> Path:
return self.project_root / ".noma"
@property
def state_file(self) -> Path:
return self.noma_dir / "state.json"
@property
def index_db(self) -> Path:
return self.noma_dir / "index.db"
# v5.1 引入: alias_index_file 已废弃,别名存储在 index.db aliases 表
@property
def chapters_dir(self) -> Path:
return self.project_root / "正文"
@property
def settings_dir(self) -> Path:
return self.project_root / "设定集"
@property
def outline_dir(self) -> Path:
return self.project_root / "大纲"
@property
def wiki_dir(self) -> Path:
return self.noma_dir / "wiki"
# ================= Embedding API 配置 =================
embed_api_type: str = "openai"
embed_base_url: str = field(default_factory=lambda: os.getenv("EMBED_BASE_URL", "https://api-inference.modelscope.cn/v1"))
embed_model: str = field(default_factory=lambda: os.getenv("EMBED_MODEL", "Qwen/Qwen3-Embedding-8B"))
embed_api_key: str = field(default_factory=lambda: os.getenv("EMBED_API_KEY", ""))
@property
def embed_url(self) -> str:
return self.embed_base_url
# ================= Rerank API 配置 =================
rerank_api_type: str = "openai"
rerank_base_url: str = field(default_factory=lambda: os.getenv("RERANK_BASE_URL", "https://api.jina.ai/v1"))
rerank_model: str = field(default_factory=lambda: os.getenv("RERANK_MODEL", "jina-reranker-v3"))
rerank_api_key: str = field(default_factory=lambda: os.getenv("RERANK_API_KEY", ""))
@property
def rerank_url(self) -> str:
return self.rerank_base_url
# ================= 并发配置 =================
embed_concurrency: int = 64
rerank_concurrency: int = 32
embed_batch_size: int = 64
# ================= 超时配置 =================
cold_start_timeout: int = 300
normal_timeout: int = 180
# ================= 重试配置 =================
api_max_retries: int = 3 # 最大重试次数
api_retry_delay: float = 1.0 # 初始重试延迟(秒),使用指数退避
# ================= 检索配置 =================
vector_top_k: int = 30
bm25_top_k: int = 20
rerank_top_n: int = 10
rrf_k: int = 60
vector_full_scan_max_vectors: int = 500
vector_prefilter_bm25_candidates: int = 200
vector_prefilter_recent_candidates: int = 200
# ================= Graph-RAG 配置 =================
graph_rag_enabled: bool = False
graph_rag_expand_hops: int = 1
graph_rag_max_expanded_entities: int = 30
graph_rag_candidate_limit: int = 150
graph_rag_boost_same_entity: float = 0.2
graph_rag_boost_related_entity: float = 0.1
graph_rag_boost_recency: float = 0.05
relationship_graph_from_index_enabled: bool = True
# ================= 实体提取配置 =================
extraction_confidence_high: float = 0.8
extraction_confidence_medium: float = 0.5
# ================= 列表截断限制 =================
max_disambiguation_warnings: int = 500
max_disambiguation_pending: int = 1000
max_state_changes: int = 2000
context_recent_summaries_window: int = 3
context_recent_meta_window: int = 3
context_alerts_slice: int = 10
context_max_appearing_characters: int = 10
context_max_urgent_foreshadowing: int = 5
context_story_skeleton_interval: int = 20
context_story_skeleton_max_samples: int = 5
context_story_skeleton_snippet_chars: int = 400
context_extra_section_budget: int = 800
context_ranker_enabled: bool = True
context_ranker_recency_weight: float = 0.7
context_ranker_frequency_weight: float = 0.3
context_ranker_hook_bonus: float = 0.2
context_ranker_length_bonus_cap: float = 0.2
context_ranker_alert_critical_keywords: tuple[str, ...] = (
"冲突",
"矛盾",
"critical",
"break",
"违规",
"断裂",
)
context_ranker_debug: bool = False
context_reader_signal_enabled: bool = True
context_reader_signal_recent_limit: int = 5
context_reader_signal_window_chapters: int = 20
context_reader_signal_review_window: int = 5
context_reader_signal_include_debt: bool = False
context_genre_profile_enabled: bool = True
context_genre_profile_max_refs: int = 8
context_genre_profile_fallback: str = "shuangwen"
context_compact_text_enabled: bool = True
context_compact_min_budget: int = 120
context_compact_head_ratio: float = 0.65
context_writing_guidance_enabled: bool = True
context_writing_guidance_max_items: int = 6
context_writing_guidance_low_score_threshold: float = 75.0
context_writing_guidance_hook_diversify: bool = True
context_methodology_enabled: bool = True
context_methodology_genre_whitelist: tuple[str, ...] = ("*",)
context_methodology_label: str = "digital-serial-v1"
context_writing_checklist_enabled: bool = True
context_writing_checklist_min_items: int = 3
context_writing_checklist_max_items: int = 6
context_writing_checklist_default_weight: float = 1.0
context_writing_score_persist_enabled: bool = True
context_writing_score_include_reader_trend: bool = True
context_writing_score_trend_window: int = 10
context_rag_assist_enabled: bool = True
context_rag_assist_top_k: int = 4
context_rag_assist_min_outline_chars: int = 40
context_rag_assist_max_query_chars: int = 120
context_dynamic_budget_enabled: bool = True
context_dynamic_budget_early_chapter: int = 30
context_dynamic_budget_late_chapter: int = 120
context_dynamic_budget_early_core_bonus: float = 0.08
context_dynamic_budget_early_scene_bonus: float = 0.04
context_dynamic_budget_late_global_bonus: float = 0.08
context_dynamic_budget_late_scene_penalty: float = 0.06
context_template_weights_dynamic: dict[str, dict[str, dict[str, float]]] = field(
default_factory=_default_context_template_weights_dynamic
)
context_genre_profile_support_composite: bool = True
context_genre_profile_max_genres: int = 2
context_genre_profile_separators: tuple[str, ...] = (
"+",
"/",
"|",
",",
",",
"、",
)
export_recent_changes_slice: int = 20
export_disambiguation_slice: int = 20
# ================= 查询默认限制 =================
query_recent_chapters_limit: int = 10
query_scenes_by_location_limit: int = 20
query_entity_appearances_limit: int = 50
query_recent_appearances_limit: int = 20
# ================= 伏笔紧急度 =================
foreshadowing_urgency_pending_high: int = 100
foreshadowing_urgency_pending_medium: int = 50
foreshadowing_urgency_target_proximity: int = 5
foreshadowing_urgency_score_high: int = 100
foreshadowing_urgency_score_medium: int = 60
foreshadowing_urgency_score_target: int = 80
foreshadowing_urgency_score_low: int = 20
foreshadowing_urgency_threshold_show: int = 60
foreshadowing_tier_weight_core: float = 3.0
foreshadowing_tier_weight_sub: float = 2.0
foreshadowing_tier_weight_decor: float = 1.0
# ================= 角色活跃度 =================
character_absence_warning: int = 30
character_absence_critical: int = 100
character_candidates_limit: int = 800
# ================= Strand Weave 节奏 =================
strand_quest_max_consecutive: int = 5
strand_fire_max_gap: int = 10
strand_constellation_max_gap: int = 15
strand_quest_ratio_min: int = 55
strand_quest_ratio_max: int = 65
strand_fire_ratio_min: int = 20
strand_fire_ratio_max: int = 30
strand_constellation_ratio_min: int = 10
strand_constellation_ratio_max: int = 20
# ================= 爽点节奏 =================
pacing_segment_size: int = 100
pacing_words_per_point_excellent: int = 1000
pacing_words_per_point_good: int = 1500
pacing_words_per_point_acceptable: int = 2000
# ================= RAG 存储 =================
@property
def rag_db(self) -> Path:
return self.noma_dir / "rag.db"
@property
def vector_db(self) -> Path:
return self.noma_dir / "vectors.db"
def ensure_dirs(self):
self.noma_dir.mkdir(parents=True, exist_ok=True)
@classmethod
def from_project_root(cls, project_root: str | Path) -> "DataModulesConfig":
root = normalize_windows_path(project_root).expanduser().resolve()
# 在构造配置前加载项目级 `.env`,以确保 EMBED_*/RERANK_* 等字段可生效
_load_project_dotenv(root)
return cls(project_root=root)
_default_config: Optional[DataModulesConfig] = None
def get_config(project_root: Optional[Path] = None) -> DataModulesConfig:
global _default_config
if project_root is not None:
return DataModulesConfig.from_project_root(project_root)
if _default_config is None:
# 默认不要盲目以 CWD 作为 project_root(很容易写到错误目录)。
# 使用统一的 project_locator 自动探测:
# - 支持 NOMA_PROJECT_ROOT
# - 支持 `.claude/.noma-current-project` 指针文件
# - 支持从当前目录/父目录寻找 `.noma/state.json`
from project_locator import resolve_project_root
root = resolve_project_root()
_default_config = DataModulesConfig.from_project_root(root)
return _default_config
def set_project_root(project_root: str | Path):
global _default_config
_default_config = DataModulesConfig.from_project_root(project_root)
@@ -0,0 +1,809 @@
#!/usr/bin/env python3
# -*- coding: utf-8 -*-
"""
ContextManager - assemble context packs with weighted priorities.
"""
from __future__ import annotations
import json
import re
import sys
import logging
from pathlib import Path
from runtime_compat import enable_windows_utf8_stdio
from typing import Any, Dict, List, Optional
try:
from chapter_outline_loader import load_chapter_outline
except ImportError: # pragma: no cover
from scripts.chapter_outline_loader import load_chapter_outline
from .config import get_config
from .index_manager import IndexManager, WritingChecklistScoreMeta
from .context_ranker import ContextRanker
from .snapshot_manager import SnapshotManager, SnapshotVersionMismatch
from .context_weights import (
DEFAULT_TEMPLATE as CONTEXT_DEFAULT_TEMPLATE,
TEMPLATE_WEIGHTS as CONTEXT_TEMPLATE_WEIGHTS,
TEMPLATE_WEIGHTS_DYNAMIC_DEFAULT as CONTEXT_TEMPLATE_WEIGHTS_DYNAMIC_DEFAULT,
)
from .genre_aliases import normalize_genre_token, to_profile_key
from .genre_profile_builder import (
build_composite_genre_hints,
extract_genre_section,
extract_markdown_refs,
parse_genre_tokens,
)
from .writing_guidance_builder import (
build_methodology_guidance_items,
build_methodology_strategy_card,
build_guidance_items,
build_writing_checklist,
is_checklist_item_completed,
)
logger = logging.getLogger(__name__)
class ContextManager:
DEFAULT_TEMPLATE = CONTEXT_DEFAULT_TEMPLATE
TEMPLATE_WEIGHTS = CONTEXT_TEMPLATE_WEIGHTS
TEMPLATE_WEIGHTS_DYNAMIC = CONTEXT_TEMPLATE_WEIGHTS_DYNAMIC_DEFAULT
EXTRA_SECTIONS = {
"story_skeleton",
"memory",
"preferences",
"alerts",
"reader_signal",
"genre_profile",
"writing_guidance",
"wiki",
}
SECTION_ORDER = [
"core",
"scene",
"global",
"reader_signal",
"genre_profile",
"writing_guidance",
"story_skeleton",
"memory",
"wiki",
"preferences",
"alerts",
]
SUMMARY_SECTION_RE = re.compile(r"##\s*剧情摘要\s*\r?\n(.*?)(?=\r?\n##|\Z)", re.DOTALL)
def __init__(self, config=None, snapshot_manager: Optional[SnapshotManager] = None):
self.config = config or get_config()
self.snapshot_manager = snapshot_manager or SnapshotManager(self.config)
self.index_manager = IndexManager(self.config)
self.context_ranker = ContextRanker(self.config)
def _is_snapshot_compatible(self, cached: Dict[str, Any], template: str) -> bool:
"""判断快照是否可用于当前模板。"""
if not isinstance(cached, dict):
return False
meta = cached.get("meta")
if not isinstance(meta, dict):
# 兼容旧快照:未记录 template 时仅允许默认模板复用
return template == self.DEFAULT_TEMPLATE
cached_template = meta.get("template")
if not isinstance(cached_template, str):
return template == self.DEFAULT_TEMPLATE
return cached_template == template
def build_context(
self,
chapter: int,
template: str | None = None,
use_snapshot: bool = True,
save_snapshot: bool = True,
max_chars: Optional[int] = None,
) -> Dict[str, Any]:
template = template or self.DEFAULT_TEMPLATE
self._active_template = template
if template not in self.TEMPLATE_WEIGHTS:
template = self.DEFAULT_TEMPLATE
self._active_template = template
if use_snapshot:
try:
cached = self.snapshot_manager.load_snapshot(chapter)
if cached and self._is_snapshot_compatible(cached, template):
return cached.get("payload", cached)
except SnapshotVersionMismatch:
# Snapshot incompatible; rebuild below.
pass
pack = self._build_pack(chapter)
if getattr(self.config, "context_ranker_enabled", True):
pack = self.context_ranker.rank_pack(pack, chapter)
assembled = self.assemble_context(pack, template=template, max_chars=max_chars)
if save_snapshot:
meta = {"template": template}
self.snapshot_manager.save_snapshot(chapter, assembled, meta=meta)
return assembled
def assemble_context(
self,
pack: Dict[str, Any],
template: str = DEFAULT_TEMPLATE,
max_chars: Optional[int] = None,
) -> Dict[str, Any]:
chapter = int((pack.get("meta") or {}).get("chapter") or 0)
weights = self._resolve_template_weights(template=template, chapter=chapter)
max_chars = max_chars or 8000
extra_budget = int(self.config.context_extra_section_budget or 0)
sections = {}
for section_name in self.SECTION_ORDER:
if section_name in pack:
sections[section_name] = pack[section_name]
assembled: Dict[str, Any] = {"meta": pack.get("meta", {}), "sections": {}}
for name, content in sections.items():
weight = weights.get(name, 0.0)
if weight > 0:
budget = int(max_chars * weight)
elif name in self.EXTRA_SECTIONS and extra_budget > 0:
budget = extra_budget
else:
budget = None
text = self._compact_json_text(content, budget)
assembled["sections"][name] = {"content": content, "text": text, "budget": budget}
assembled["template"] = template
assembled["weights"] = weights
if chapter > 0:
assembled.setdefault("meta", {})["context_weight_stage"] = self._resolve_context_stage(chapter)
return assembled
def filter_invalid_items(self, items: List[Dict[str, Any]], source_type: str, id_key: str) -> List[Dict[str, Any]]:
confirmed = self.index_manager.get_invalid_ids(source_type, status="confirmed")
pending = self.index_manager.get_invalid_ids(source_type, status="pending")
result = []
for item in items:
item_id = str(item.get(id_key, ""))
if item_id in confirmed:
continue
if item_id in pending:
item = dict(item)
item["warning"] = "pending_invalid"
result.append(item)
return result
def apply_confidence_filter(self, items: List[Dict[str, Any]], min_confidence: float) -> List[Dict[str, Any]]:
filtered: List[Dict[str, Any]] = []
for item in items:
conf = item.get("confidence")
if conf is None or conf >= min_confidence:
filtered.append(item)
return filtered
def _build_pack(self, chapter: int) -> Dict[str, Any]:
state = self._load_state()
core = {
"chapter_outline": self._load_outline(chapter),
"protagonist_snapshot": state.get("protagonist_state", {}),
"recent_summaries": self._load_recent_summaries(
chapter,
window=self.config.context_recent_summaries_window,
),
"recent_meta": self._load_recent_meta(
state,
chapter,
window=self.config.context_recent_meta_window,
),
}
scene = {
"location_context": state.get("protagonist_state", {}).get("location", {}),
"appearing_characters": self._load_recent_appearances(
limit=self.config.context_max_appearing_characters,
),
}
scene["appearing_characters"] = self.filter_invalid_items(
scene["appearing_characters"], source_type="entity", id_key="entity_id"
)
global_ctx = {
"worldview_skeleton": self._load_setting("世界观"),
"power_system_skeleton": self._load_setting("力量体系"),
"style_contract_ref": self._load_setting("风格契约"),
}
preferences = self._load_json_optional(self.config.noma_dir / "preferences.json")
memory = self._load_json_optional(self.config.noma_dir / "project_memory.json")
story_skeleton = self._load_story_skeleton(chapter)
alert_slice = max(0, int(self.config.context_alerts_slice))
reader_signal = self._load_reader_signal(chapter)
genre_profile = self._load_genre_profile(state)
writing_guidance = self._build_writing_guidance(chapter, reader_signal, genre_profile)
wiki_data = self._load_wiki_context()
return {
"meta": {"chapter": chapter},
"core": core,
"scene": scene,
"global": global_ctx,
"reader_signal": reader_signal,
"genre_profile": genre_profile,
"writing_guidance": writing_guidance,
"story_skeleton": story_skeleton,
"preferences": preferences,
"memory": memory,
"wiki": wiki_data,
"alerts": {
"disambiguation_warnings": (
state.get("disambiguation_warnings", [])[-alert_slice:] if alert_slice else []
),
"disambiguation_pending": (
state.get("disambiguation_pending", [])[-alert_slice:] if alert_slice else []
),
},
}
def _load_wiki_context(self) -> Dict[str, Any]:
"""Load relevant wiki entries for context assembly."""
from .wiki_manager import WikiManager
wiki_dir = self.config.wiki_dir
if not wiki_dir.exists():
return {}
wiki = WikiManager(self.config)
protagonist_wiki = None
protagonist = self.index_manager.get_protagonist()
if protagonist:
eid = protagonist.get("id", "")
if eid:
entry = wiki.get_entity_wiki(eid)
if entry:
protagonist_wiki = entry.get("frontmatter", {})
plot_threads = wiki.get_plot_threads()
patterns = wiki.get_writing_patterns()
return {
"protagonist_profile": protagonist_wiki,
"plot_threads": plot_threads.get("body", "")[:500] if plot_threads else None,
"writing_patterns": patterns[-5:] if patterns else [],
}
def _load_reader_signal(self, chapter: int) -> Dict[str, Any]:
if not getattr(self.config, "context_reader_signal_enabled", True):
return {}
recent_limit = max(1, int(getattr(self.config, "context_reader_signal_recent_limit", 5)))
pattern_window = max(1, int(getattr(self.config, "context_reader_signal_window_chapters", 20)))
review_window = max(1, int(getattr(self.config, "context_reader_signal_review_window", 5)))
include_debt = bool(getattr(self.config, "context_reader_signal_include_debt", False))
recent_power = self.index_manager.get_recent_reading_power(limit=recent_limit)
pattern_stats = self.index_manager.get_pattern_usage_stats(last_n_chapters=pattern_window)
hook_stats = self.index_manager.get_hook_type_stats(last_n_chapters=pattern_window)
review_trend = self.index_manager.get_review_trend_stats(last_n=review_window)
low_score_ranges: List[Dict[str, Any]] = []
for row in review_trend.get("recent_ranges", []):
score = row.get("overall_score")
if isinstance(score, (int, float)) and float(score) < 75:
low_score_ranges.append(
{
"start_chapter": row.get("start_chapter"),
"end_chapter": row.get("end_chapter"),
"overall_score": score,
}
)
signal: Dict[str, Any] = {
"recent_reading_power": recent_power,
"pattern_usage": pattern_stats,
"hook_type_usage": hook_stats,
"review_trend": review_trend,
"low_score_ranges": low_score_ranges,
"next_chapter": chapter,
}
if include_debt:
signal["debt_summary"] = self.index_manager.get_debt_summary()
return signal
def _load_genre_profile(self, state: Dict[str, Any]) -> Dict[str, Any]:
if not getattr(self.config, "context_genre_profile_enabled", True):
return {}
fallback = str(getattr(self.config, "context_genre_profile_fallback", "shuangwen") or "shuangwen")
project = state.get("project") or {}
project_info = state.get("project_info") or {}
genre_raw = str(project.get("genre") or project_info.get("genre") or fallback)
genres = self._parse_genre_tokens(genre_raw)
if not genres:
genres = [fallback]
max_genres = max(1, int(getattr(self.config, "context_genre_profile_max_genres", 2)))
genres = genres[:max_genres]
primary_genre = genres[0]
secondary_genres = genres[1:]
composite = len(genres) > 1
profile_path = self.config.project_root / ".claude" / "references" / "genre-profiles.md"
taxonomy_path = self.config.project_root / ".claude" / "references" / "reading-power-taxonomy.md"
profile_text = profile_path.read_text(encoding="utf-8") if profile_path.exists() else ""
taxonomy_text = taxonomy_path.read_text(encoding="utf-8") if taxonomy_path.exists() else ""
profile_excerpt = self._extract_genre_section(profile_text, primary_genre)
taxonomy_excerpt = self._extract_genre_section(taxonomy_text, primary_genre)
secondary_profiles: List[str] = []
secondary_taxonomies: List[str] = []
for extra in secondary_genres:
secondary_profiles.append(self._extract_genre_section(profile_text, extra))
secondary_taxonomies.append(self._extract_genre_section(taxonomy_text, extra))
refs = self._extract_markdown_refs(
"\n".join([profile_excerpt] + secondary_profiles),
max_items=int(getattr(self.config, "context_genre_profile_max_refs", 8)),
)
composite_hints = self._build_composite_genre_hints(genres, refs)
return {
"genre": primary_genre,
"genre_raw": genre_raw,
"genres": genres,
"composite": composite,
"secondary_genres": secondary_genres,
"profile_excerpt": profile_excerpt,
"taxonomy_excerpt": taxonomy_excerpt,
"secondary_profile_excerpts": secondary_profiles,
"secondary_taxonomy_excerpts": secondary_taxonomies,
"reference_hints": refs,
"composite_hints": composite_hints,
}
def _build_writing_guidance(
self,
chapter: int,
reader_signal: Dict[str, Any],
genre_profile: Dict[str, Any],
) -> Dict[str, Any]:
if not getattr(self.config, "context_writing_guidance_enabled", True):
return {}
limit = max(1, int(getattr(self.config, "context_writing_guidance_max_items", 6)))
low_score_threshold = float(
getattr(self.config, "context_writing_guidance_low_score_threshold", 75.0)
)
guidance_bundle = build_guidance_items(
chapter=chapter,
reader_signal=reader_signal,
genre_profile=genre_profile,
low_score_threshold=low_score_threshold,
hook_diversify_enabled=bool(
getattr(self.config, "context_writing_guidance_hook_diversify", True)
),
)
guidance = list(guidance_bundle.get("guidance") or [])
methodology_strategy: Dict[str, Any] = {}
if self._is_methodology_enabled_for_genre(genre_profile):
methodology_strategy = build_methodology_strategy_card(
chapter=chapter,
reader_signal=reader_signal,
genre_profile=genre_profile,
label=str(getattr(self.config, "context_methodology_label", "digital-serial-v1")),
)
guidance.extend(build_methodology_guidance_items(methodology_strategy))
checklist = self._build_writing_checklist(
chapter=chapter,
guidance_items=guidance,
reader_signal=reader_signal,
genre_profile=genre_profile,
strategy_card=methodology_strategy,
)
checklist_score = self._compute_writing_checklist_score(
chapter=chapter,
checklist=checklist,
reader_signal=reader_signal,
)
if getattr(self.config, "context_writing_score_persist_enabled", True):
self._persist_writing_checklist_score(checklist_score)
low_ranges = guidance_bundle.get("low_ranges") or []
hook_usage = guidance_bundle.get("hook_usage") or {}
pattern_usage = guidance_bundle.get("pattern_usage") or {}
genre = str(guidance_bundle.get("genre") or genre_profile.get("genre") or "").strip()
hook_types = list(hook_usage.keys())[:3] if isinstance(hook_usage, dict) else []
top_patterns = (
sorted(pattern_usage, key=pattern_usage.get, reverse=True)[:3]
if isinstance(pattern_usage, dict)
else []
)
return {
"chapter": chapter,
"guidance_items": guidance[:limit],
"checklist": checklist,
"checklist_score": checklist_score,
"methodology": methodology_strategy,
"signals_used": {
"has_low_score_ranges": bool(low_ranges),
"hook_types": hook_types,
"top_patterns": top_patterns,
"genre": genre,
"methodology_enabled": bool(methodology_strategy.get("enabled")),
},
}
def _compute_writing_checklist_score(
self,
chapter: int,
checklist: List[Dict[str, Any]],
reader_signal: Dict[str, Any],
) -> Dict[str, Any]:
total_items = len(checklist)
required_items = 0
completed_items = 0
completed_required = 0
total_weight = 0.0
completed_weight = 0.0
pending_labels: List[str] = []
for item in checklist:
if not isinstance(item, dict):
continue
required = bool(item.get("required"))
weight = float(item.get("weight") or 1.0)
total_weight += weight
if required:
required_items += 1
completed = self._is_checklist_item_completed(item, reader_signal)
if completed:
completed_items += 1
completed_weight += weight
if required:
completed_required += 1
else:
pending_labels.append(str(item.get("label") or item.get("id") or "未命名项"))
completion_rate = (completed_items / total_items) if total_items > 0 else 1.0
weighted_rate = (completed_weight / total_weight) if total_weight > 0 else completion_rate
required_rate = (completed_required / required_items) if required_items > 0 else 1.0
score = 100.0 * (0.5 * weighted_rate + 0.3 * required_rate + 0.2 * completion_rate)
if getattr(self.config, "context_writing_score_include_reader_trend", True):
trend_window = max(1, int(getattr(self.config, "context_writing_score_trend_window", 10)))
trend = self.index_manager.get_writing_checklist_score_trend(last_n=trend_window)
baseline = float(trend.get("score_avg") or 0.0)
if baseline > 0:
score += max(-10.0, min(10.0, (score - baseline) * 0.1))
score = round(max(0.0, min(100.0, score)), 2)
return {
"chapter": chapter,
"score": score,
"completion_rate": round(completion_rate, 4),
"weighted_completion_rate": round(weighted_rate, 4),
"required_completion_rate": round(required_rate, 4),
"total_items": total_items,
"required_items": required_items,
"completed_items": completed_items,
"completed_required": completed_required,
"total_weight": round(total_weight, 2),
"completed_weight": round(completed_weight, 2),
"pending_items": pending_labels,
"trend_window": int(getattr(self.config, "context_writing_score_trend_window", 10)),
}
def _is_checklist_item_completed(self, item: Dict[str, Any], reader_signal: Dict[str, Any]) -> bool:
return is_checklist_item_completed(item, reader_signal)
def _persist_writing_checklist_score(self, checklist_score: Dict[str, Any]) -> None:
if not checklist_score:
return
try:
self.index_manager.save_writing_checklist_score(
WritingChecklistScoreMeta(
chapter=int(checklist_score.get("chapter") or 0),
template=str(getattr(self, "_active_template", self.DEFAULT_TEMPLATE) or self.DEFAULT_TEMPLATE),
total_items=int(checklist_score.get("total_items") or 0),
required_items=int(checklist_score.get("required_items") or 0),
completed_items=int(checklist_score.get("completed_items") or 0),
completed_required=int(checklist_score.get("completed_required") or 0),
total_weight=float(checklist_score.get("total_weight") or 0.0),
completed_weight=float(checklist_score.get("completed_weight") or 0.0),
completion_rate=float(checklist_score.get("completion_rate") or 0.0),
score=float(checklist_score.get("score") or 0.0),
score_breakdown={
"weighted_completion_rate": checklist_score.get("weighted_completion_rate"),
"required_completion_rate": checklist_score.get("required_completion_rate"),
"trend_window": checklist_score.get("trend_window"),
},
pending_items=list(checklist_score.get("pending_items") or []),
source="context_manager",
)
)
except Exception as exc:
logger.warning("failed to persist writing checklist score: %s", exc)
def _resolve_context_stage(self, chapter: int) -> str:
early = max(1, int(getattr(self.config, "context_dynamic_budget_early_chapter", 30)))
late = max(early + 1, int(getattr(self.config, "context_dynamic_budget_late_chapter", 120)))
if chapter <= early:
return "early"
if chapter >= late:
return "late"
return "mid"
def _resolve_template_weights(self, template: str, chapter: int) -> Dict[str, float]:
template_key = template if template in self.TEMPLATE_WEIGHTS else self.DEFAULT_TEMPLATE
base = dict(self.TEMPLATE_WEIGHTS.get(template_key, self.TEMPLATE_WEIGHTS[self.DEFAULT_TEMPLATE]))
if not getattr(self.config, "context_dynamic_budget_enabled", True):
return base
stage = self._resolve_context_stage(chapter)
dynamic_weights = getattr(self.config, "context_template_weights_dynamic", None)
if not isinstance(dynamic_weights, dict):
dynamic_weights = self.TEMPLATE_WEIGHTS_DYNAMIC
stage_weights = dynamic_weights.get(stage, {}) if isinstance(dynamic_weights.get(stage, {}), dict) else {}
staged = stage_weights.get(template_key)
if isinstance(staged, dict):
return dict(staged)
return base
def _parse_genre_tokens(self, genre_raw: str) -> List[str]:
support_composite = bool(getattr(self.config, "context_genre_profile_support_composite", True))
separators_raw = getattr(self.config, "context_genre_profile_separators", ("+", "/", "|", ","))
separators = tuple(str(token) for token in separators_raw if str(token))
return parse_genre_tokens(
genre_raw,
support_composite=support_composite,
separators=separators,
)
def _normalize_genre_token(self, token: str) -> str:
return normalize_genre_token(token)
def _build_composite_genre_hints(self, genres: List[str], refs: List[str]) -> List[str]:
return build_composite_genre_hints(genres, refs)
def _build_writing_checklist(
self,
chapter: int,
guidance_items: List[str],
reader_signal: Dict[str, Any],
genre_profile: Dict[str, Any],
strategy_card: Dict[str, Any] | None = None,
) -> List[Dict[str, Any]]:
_ = chapter
if not getattr(self.config, "context_writing_checklist_enabled", True):
return []
min_items = max(1, int(getattr(self.config, "context_writing_checklist_min_items", 3)))
max_items = max(min_items, int(getattr(self.config, "context_writing_checklist_max_items", 6)))
default_weight = float(getattr(self.config, "context_writing_checklist_default_weight", 1.0))
if default_weight <= 0:
default_weight = 1.0
return build_writing_checklist(
guidance_items=guidance_items,
reader_signal=reader_signal,
genre_profile=genre_profile,
strategy_card=strategy_card,
min_items=min_items,
max_items=max_items,
default_weight=default_weight,
)
def _is_methodology_enabled_for_genre(self, genre_profile: Dict[str, Any]) -> bool:
if not bool(getattr(self.config, "context_methodology_enabled", False)):
return False
whitelist_raw = getattr(self.config, "context_methodology_genre_whitelist", ("*",))
if isinstance(whitelist_raw, str):
whitelist_iter = [whitelist_raw]
else:
whitelist_iter = list(whitelist_raw or [])
whitelist = {str(token).strip().lower() for token in whitelist_iter if str(token).strip()}
if not whitelist:
return True
if "*" in whitelist or "all" in whitelist:
return True
genre = str((genre_profile or {}).get("genre") or "").strip()
if not genre:
return False
profile_key = to_profile_key(genre)
return profile_key in whitelist
def _compact_json_text(self, content: Any, budget: Optional[int]) -> str:
raw = json.dumps(content, ensure_ascii=False)
if budget is None or len(raw) <= budget:
return raw
if not getattr(self.config, "context_compact_text_enabled", True):
return raw[:budget]
min_budget = max(1, int(getattr(self.config, "context_compact_min_budget", 120)))
if budget <= min_budget:
return raw[:budget]
head_ratio = float(getattr(self.config, "context_compact_head_ratio", 0.65))
head_budget = int(budget * max(0.2, min(0.9, head_ratio)))
tail_budget = max(0, budget - head_budget - 10)
compact = f"{raw[:head_budget]}…[TRUNCATED]{raw[-tail_budget:] if tail_budget else ''}"
return compact[:budget]
def _extract_genre_section(self, text: str, genre: str) -> str:
return extract_genre_section(text, genre)
def _extract_markdown_refs(self, text: str, max_items: int = 8) -> List[str]:
return extract_markdown_refs(text, max_items=max_items)
def _load_state(self) -> Dict[str, Any]:
path = self.config.state_file
if not path.exists():
return {}
return json.loads(path.read_text(encoding="utf-8"))
def _load_outline(self, chapter: int) -> str:
return load_chapter_outline(self.config.project_root, chapter, max_chars=1500)
def _load_recent_summaries(self, chapter: int, window: int = 3) -> List[Dict[str, Any]]:
summaries = []
for ch in range(max(1, chapter - window), chapter):
summary = self._load_summary_text(ch)
if summary:
summaries.append(summary)
return summaries
def _load_recent_meta(self, state: Dict[str, Any], chapter: int, window: int = 3) -> List[Dict[str, Any]]:
meta = state.get("chapter_meta", {}) or {}
results = []
for ch in range(max(1, chapter - window), chapter):
for key in (f"{ch:04d}", str(ch)):
if key in meta:
results.append({"chapter": ch, **meta.get(key, {})})
break
return results
def _load_recent_appearances(self, limit: Optional[int] = None) -> List[Dict[str, Any]]:
appearances = self.index_manager.get_recent_appearances(limit=limit)
return appearances or []
def _load_setting(self, keyword: str) -> str:
settings_dir = self.config.settings_dir
candidates = [
settings_dir / f"{keyword}.md",
]
for path in candidates:
if path.exists():
return path.read_text(encoding="utf-8")
# fallback: any file containing keyword
matches = list(settings_dir.glob(f"*{keyword}*.md"))
if matches:
return matches[0].read_text(encoding="utf-8")
return f"[{keyword}设定未找到]"
def _extract_summary_excerpt(self, text: str, max_chars: int) -> str:
if not text:
return ""
match = self.SUMMARY_SECTION_RE.search(text)
excerpt = match.group(1).strip() if match else text.strip()
if max_chars > 0 and len(excerpt) > max_chars:
return excerpt[:max_chars].rstrip()
return excerpt
def _load_summary_text(self, chapter: int, snippet_chars: Optional[int] = None) -> Optional[Dict[str, Any]]:
summary_path = self.config.noma_dir / "summaries" / f"ch{chapter:04d}.md"
if not summary_path.exists():
return None
text = summary_path.read_text(encoding="utf-8")
if snippet_chars:
summary_text = self._extract_summary_excerpt(text, snippet_chars)
else:
summary_text = text
return {"chapter": chapter, "summary": summary_text}
def _load_story_skeleton(self, chapter: int) -> List[Dict[str, Any]]:
interval = max(1, int(self.config.context_story_skeleton_interval))
max_samples = max(0, int(self.config.context_story_skeleton_max_samples))
snippet_chars = int(self.config.context_story_skeleton_snippet_chars)
if max_samples <= 0 or chapter <= interval:
return []
samples: List[Dict[str, Any]] = []
cursor = chapter - interval
while cursor >= 1 and len(samples) < max_samples:
summary = self._load_summary_text(cursor, snippet_chars=snippet_chars)
if summary and summary.get("summary"):
samples.append(summary)
cursor -= interval
samples.reverse()
return samples
def _load_json_optional(self, path: Path) -> Dict[str, Any]:
if not path.exists():
return {}
try:
return json.loads(path.read_text(encoding="utf-8"))
except json.JSONDecodeError:
return {}
def main():
import argparse
from .cli_output import print_success, print_error
parser = argparse.ArgumentParser(description="Context Manager CLI")
parser.add_argument("--project-root", type=str, help="项目根目录")
parser.add_argument("--chapter", type=int, required=True)
parser.add_argument("--template", type=str, default=ContextManager.DEFAULT_TEMPLATE)
parser.add_argument("--no-snapshot", action="store_true")
parser.add_argument("--max-chars", type=int, default=8000)
args = parser.parse_args()
config = None
if args.project_root:
# 允许传入“工作区根目录”,统一解析到真正的 book project_root(必须包含 .noma/state.json)
from project_locator import resolve_project_root
from .config import DataModulesConfig
resolved_root = resolve_project_root(args.project_root)
config = DataModulesConfig.from_project_root(resolved_root)
manager = ContextManager(config)
try:
payload = manager.build_context(
chapter=args.chapter,
template=args.template,
use_snapshot=not args.no_snapshot,
save_snapshot=True,
max_chars=args.max_chars,
)
print_success(payload, message="context_built")
try:
manager.index_manager.log_tool_call("context_manager:build", True, chapter=args.chapter)
except Exception as exc:
logger.warning("failed to log successful tool call: %s", exc)
except Exception as exc:
print_error("CONTEXT_BUILD_FAILED", str(exc), suggestion="请检查项目结构与依赖文件")
try:
manager.index_manager.log_tool_call(
"context_manager:build", False, error_code="CONTEXT_BUILD_FAILED", error_message=str(exc), chapter=args.chapter
)
except Exception as log_exc:
logger.warning("failed to log failed tool call: %s", log_exc)
if __name__ == "__main__":
import sys
if sys.platform == "win32":
enable_windows_utf8_stdio()
main()
+210
View File
@@ -0,0 +1,210 @@
#!/usr/bin/env python3
# -*- coding: utf-8 -*-
"""
Context ranker for Context Contract v2.
Goals:
- Prefer recency while keeping frequent entities stable.
- Prioritize high-signal hook/alert items.
- Keep output shape backward compatible (same keys, re-ordered lists).
"""
from __future__ import annotations
import math
from typing import Any, Dict, List, Optional
from .config import get_config
class ContextRanker:
"""Rank context-pack sections with lightweight deterministic heuristics."""
SUMMARY_HOOK_HINTS = ("?", "?", "悬念", "钩子", "反转", "冲突")
def __init__(self, config=None):
self.config = config or get_config()
def rank_pack(self, pack: Dict[str, Any], chapter: int) -> Dict[str, Any]:
ranked = dict(pack)
core = dict(ranked.get("core") or {})
core["recent_summaries"] = self.rank_recent_summaries(core.get("recent_summaries") or [], chapter)
core["recent_meta"] = self.rank_recent_meta(core.get("recent_meta") or [], chapter)
ranked["core"] = core
scene = dict(ranked.get("scene") or {})
scene["appearing_characters"] = self.rank_appearances(scene.get("appearing_characters") or [], chapter)
ranked["scene"] = scene
ranked["story_skeleton"] = self.rank_story_skeleton(ranked.get("story_skeleton") or [], chapter)
alerts = dict(ranked.get("alerts") or {})
alerts["disambiguation_warnings"] = self.rank_alerts(alerts.get("disambiguation_warnings") or [], chapter)
alerts["disambiguation_pending"] = self.rank_alerts(alerts.get("disambiguation_pending") or [], chapter)
ranked["alerts"] = alerts
meta = dict(ranked.get("meta") or {})
meta.setdefault("context_contract_version", "v2")
meta["ranker"] = {
"enabled": True,
"recency_weight": float(self.config.context_ranker_recency_weight),
"frequency_weight": float(self.config.context_ranker_frequency_weight),
"hook_bonus": float(self.config.context_ranker_hook_bonus),
}
ranked["meta"] = meta
return ranked
def rank_recent_summaries(self, items: List[Dict[str, Any]], current_chapter: int) -> List[Dict[str, Any]]:
scored = []
for raw in items:
item = dict(raw)
chapter = self._as_int(item.get("chapter"))
summary = str(item.get("summary") or "")
recency = self._recency_score(chapter, current_chapter)
frequency = self._length_score(summary)
hook_bonus = float(self.config.context_ranker_hook_bonus) if self._has_hook_hint(summary) else 0.0
score = self._combine_score(recency, frequency, hook_bonus)
scored.append(self._with_debug_score(item, score, recency, frequency, hook_bonus))
scored.sort(key=lambda row: row[0], reverse=True)
return [row[1] for row in scored]
def rank_recent_meta(self, items: List[Dict[str, Any]], current_chapter: int) -> List[Dict[str, Any]]:
scored = []
for raw in items:
item = dict(raw)
chapter = self._as_int(item.get("chapter"))
hook = str(item.get("hook") or "")
hook_bonus = float(self.config.context_ranker_hook_bonus) if hook else 0.0
recency = self._recency_score(chapter, current_chapter)
frequency = self._length_score(hook)
score = self._combine_score(recency, frequency, hook_bonus)
scored.append(self._with_debug_score(item, score, recency, frequency, hook_bonus))
scored.sort(key=lambda row: row[0], reverse=True)
return [row[1] for row in scored]
def rank_appearances(self, items: List[Dict[str, Any]], current_chapter: int) -> List[Dict[str, Any]]:
scored = []
for raw in items:
item = dict(raw)
last_chapter = self._as_int(item.get("last_chapter") or item.get("chapter"))
total = self._as_int(item.get("total")) or 0
warning_penalty = 0.15 if item.get("warning") else 0.0
recency = self._recency_score(last_chapter, current_chapter)
frequency = self._frequency_score(total)
score = self._combine_score(recency, frequency, 0.0) - warning_penalty
scored.append(self._with_debug_score(item, score, recency, frequency, -warning_penalty))
scored.sort(key=lambda row: row[0], reverse=True)
return [row[1] for row in scored]
def rank_story_skeleton(self, items: List[Dict[str, Any]], current_chapter: int) -> List[Dict[str, Any]]:
scored = []
for raw in items:
item = dict(raw)
chapter = self._as_int(item.get("chapter"))
summary = str(item.get("summary") or "")
recency = self._recency_score(chapter, current_chapter)
frequency = self._length_score(summary)
score = self._combine_score(recency, frequency, 0.0)
scored.append(self._with_debug_score(item, score, recency, frequency, 0.0))
scored.sort(key=lambda row: row[0], reverse=True)
return [row[1] for row in scored]
def rank_alerts(self, alerts: List[Any], current_chapter: int) -> List[Any]:
scored = []
keywords = tuple(self.config.context_ranker_alert_critical_keywords)
for raw in alerts:
if isinstance(raw, dict):
item: Any = dict(raw)
chapter = self._as_int(item.get("chapter"))
text = str(item.get("message") or item.get("content") or json_safe(item))
severity = str(item.get("severity") or "").lower()
critical_bonus = 0.3 if severity in {"critical", "high"} else 0.0
else:
item = raw
chapter = None
text = str(raw)
critical_bonus = 0.0
recency = self._recency_score(chapter, current_chapter)
keyword_bonus = 0.3 if any(word and word in text for word in keywords) else 0.0
score = recency + critical_bonus + keyword_bonus
if isinstance(item, dict):
scored.append(self._with_debug_score(item, score, recency, critical_bonus, keyword_bonus))
else:
scored.append((score, item))
scored.sort(key=lambda row: row[0], reverse=True)
return [row[1] for row in scored]
def _combine_score(self, recency: float, frequency: float, bonus: float) -> float:
return (
recency * float(self.config.context_ranker_recency_weight)
+ frequency * float(self.config.context_ranker_frequency_weight)
+ bonus
)
def _recency_score(self, source_chapter: Optional[int], current_chapter: int) -> float:
if source_chapter is None:
return 0.0
gap = max(0, int(current_chapter) - int(source_chapter))
return 1.0 / (1.0 + gap)
def _frequency_score(self, total: int) -> float:
if total <= 0:
return 0.0
# log scale to avoid over-favoring very frequent entities
return min(1.0, math.log(1.0 + float(total)) / math.log(11.0))
def _length_score(self, text: str) -> float:
if not text:
return 0.0
ratio = min(len(text) / 1200.0, 1.0)
cap = float(self.config.context_ranker_length_bonus_cap)
return ratio * cap
def _has_hook_hint(self, text: str) -> bool:
return any(token in text for token in self.SUMMARY_HOOK_HINTS)
def _as_int(self, value: Any) -> Optional[int]:
if value is None:
return None
try:
return int(value)
except (TypeError, ValueError):
return None
def _with_debug_score(
self,
item: Dict[str, Any],
score: float,
recency: float,
frequency: float,
bonus: float,
) -> tuple[float, Dict[str, Any]]:
if getattr(self.config, "context_ranker_debug", False):
item["_context_score"] = round(score, 6)
item["_context_score_detail"] = {
"recency": round(recency, 6),
"frequency": round(frequency, 6),
"bonus": round(bonus, 6),
}
return score, item
def json_safe(value: Any) -> str:
try:
import json
return json.dumps(value, ensure_ascii=False)
except Exception:
return str(value)
@@ -0,0 +1,39 @@
#!/usr/bin/env python3
# -*- coding: utf-8 -*-
"""
Centralized context template weights.
"""
from __future__ import annotations
DEFAULT_TEMPLATE = "plot"
TEMPLATE_WEIGHTS: dict[str, dict[str, float]] = {
"plot": {"core": 0.40, "scene": 0.35, "global": 0.25},
"battle": {"core": 0.35, "scene": 0.45, "global": 0.20},
"emotion": {"core": 0.45, "scene": 0.35, "global": 0.20},
"transition": {"core": 0.50, "scene": 0.25, "global": 0.25},
}
TEMPLATE_WEIGHTS_DYNAMIC_DEFAULT: dict[str, dict[str, dict[str, float]]] = {
"early": {
"plot": {"core": 0.48, "scene": 0.39, "global": 0.13},
"battle": {"core": 0.42, "scene": 0.50, "global": 0.08},
"emotion": {"core": 0.52, "scene": 0.38, "global": 0.10},
"transition": {"core": 0.56, "scene": 0.28, "global": 0.16},
},
"mid": {
"plot": {"core": 0.40, "scene": 0.35, "global": 0.25},
"battle": {"core": 0.35, "scene": 0.45, "global": 0.20},
"emotion": {"core": 0.45, "scene": 0.35, "global": 0.20},
"transition": {"core": 0.50, "scene": 0.25, "global": 0.25},
},
"late": {
"plot": {"core": 0.36, "scene": 0.29, "global": 0.35},
"battle": {"core": 0.31, "scene": 0.39, "global": 0.30},
"emotion": {"core": 0.41, "scene": 0.29, "global": 0.30},
"transition": {"core": 0.46, "scene": 0.21, "global": 0.33},
},
}
@@ -0,0 +1,726 @@
#!/usr/bin/env python3
# -*- coding: utf-8 -*-
"""
Cross-Project RAG - 跨层级 RAG 检索模块
实现三层 RAG 架构(向下继承):
1. 小说私有层 - novel/.noma/rag/
2. 工作空间共享层 - workspaces/{ws}/.noma/rag/
3. 工程共享层 - project_root/.noma/rag/
4. 插件内置层 - plugin/matrices/
读取时:小说 → 工作空间 → 工程 → 插件(向下继承)
写入时:默认写入小说层,可选择向上沉淀
"""
import json
import os
import sqlite3
import asyncio
from pathlib import Path
from dataclasses import dataclass, field
from typing import List, Dict, Any, Optional
from enum import Enum
from datetime import datetime
try:
from runtime_compat import enable_windows_utf8_stdio
except ImportError:
enable_windows_utf8_stdio = lambda: None
class RAGLayer(Enum):
"""RAG 层级 - 四层架构"""
NOVEL = "novel" # 小说私有(最高权重)
WORKSPACE = "workspace" # 工作空间共享
PROJECT = "project" # 工程目录共享
PLUGIN = "plugin" # 插件内置
@dataclass
class LearnedPattern:
"""学习到的爽点模式"""
pattern_id: str
pattern_type: str
title: str
description: str
tension_curve: List[List[float]]
catharsis_model: str
structure: Dict[str, int]
hot_spots: List[List[Any]]
style_tags: List[str]
source_project: str
source_chapter: int
learned_at: str
usage_count: int = 0
metadata: Dict[str, Any] = field(default_factory=dict)
@dataclass
class CrossProjectSearchResult:
"""检索结果"""
chunk_id: str
content: str
score: float
source_layer: RAGLayer
source_project: Optional[str]
chapter: Optional[int]
chunk_type: Optional[str]
metadata: Dict[str, Any] = field(default_factory=dict)
class CrossProjectRAG:
"""
四层 RAG 检索器
检索时自动向下继承:小说私有 > 工作空间 > 工程 > 插件
"""
# 权重:小说私有 > 工作空间 > 工程 > 插件
DEFAULT_WEIGHTS = {
RAGLayer.NOVEL: 1.0,
RAGLayer.WORKSPACE: 0.8,
RAGLayer.PROJECT: 0.6,
RAGLayer.PLUGIN: 0.3,
}
def __init__(
self,
project_root: Path,
workspace_root: Optional[Path] = None,
project_root_dir: Optional[Path] = None,
plugin_root: Optional[Path] = None,
weights: Optional[Dict[RAGLayer, float]] = None,
config: Optional[Any] = None,
):
self.project_root = Path(project_root).resolve()
# 四层根路径
self.novel_root = self.project_root
self.workspace_root = workspace_root or self._resolve_workspace_root()
self.project_root_dir = project_root_dir or self._resolve_project_root_dir()
self.plugin_root = plugin_root or self._resolve_plugin_root()
self.weights = weights or self.DEFAULT_WEIGHTS
# 加载配置
self.config = config
self._load_config()
# 初始化四层路径
self._init_paths()
# 确保目录存在
self._ensure_dirs()
def _init_paths(self):
"""初始化四层 RAG 路径"""
# 层1:小说私有
self.novel_rag_dir = self.novel_root / ".noma" / "rag"
self.novel_vectors_db = self.novel_rag_dir / "vectors.db"
self.novel_learned_db = self.novel_rag_dir / "learned.db"
# 层2:工作空间共享
if self.workspace_root:
self.ws_rag_dir = self.workspace_root / ".noma" / "rag"
self.ws_shared_dir = self.ws_rag_dir / "shared"
self.ws_catharsis_dir = self.ws_shared_dir / "catharsis"
self.ws_genres_dir = self.ws_shared_dir / "genres"
self.ws_learned_dir = self.ws_rag_dir / "learned"
else:
self.ws_rag_dir = None
self.ws_shared_dir = None
self.ws_catharsis_dir = None
self.ws_genres_dir = None
self.ws_learned_dir = None
# 层3:工程共享
if self.project_root_dir:
self.proj_rag_dir = self.project_root_dir / ".noma" / "rag"
self.proj_shared_dir = self.proj_rag_dir / "shared"
self.proj_catharsis_dir = self.proj_shared_dir / "catharsis"
self.proj_genres_dir = self.proj_shared_dir / "genres"
self.proj_learned_dir = self.proj_rag_dir / "learned"
else:
self.proj_rag_dir = None
self.proj_shared_dir = None
self.proj_catharsis_dir = None
self.proj_genres_dir = None
self.proj_learned_dir = None
# 层4:插件内置
self.plugin_matrices_dir = self.plugin_root / "matrices"
self.plugin_catharsis_dir = self.plugin_matrices_dir / "catharsis_models"
self.plugin_genres_dir = self.plugin_matrices_dir / "genres"
def _resolve_workspace_root(self) -> Optional[Path]:
"""解析工作空间根目录"""
env_path = os.environ.get("NOMA_WORKSPACE_ROOT")
if env_path:
p = Path(env_path).resolve()
if p.exists() and (p / ".noma" / "rag").exists():
return p
current = self.project_root
while True:
workspaces_dir = current / "workspaces"
if workspaces_dir.is_dir():
for ws_dir in workspaces_dir.iterdir():
if ws_dir.is_dir() and (ws_dir / ".noma" / "rag").exists():
if self.project_root == ws_dir or str(self.project_root).startswith(str(ws_dir) + os.sep):
return ws_dir
parent = current.parent
if parent == current:
break
current = parent
return None
def _resolve_project_root_dir(self) -> Optional[Path]:
"""解析工程目录根"""
env_path = os.environ.get("NOMA_PROJECT_ROOT_DIR")
if env_path:
p = Path(env_path).resolve()
if p.exists() and (p / "workspaces").is_dir() and (p / ".noma" / "rag").is_dir():
return p
current = self.project_root
while True:
if (current / "workspaces").is_dir() and (current / ".noma" / "rag").is_dir():
return current
parent = current.parent
if parent == current:
break
current = parent
return None
def _resolve_plugin_root(self) -> Path:
"""解析插件根目录"""
env_path = os.environ.get("NOMA_PLUGIN_ROOT")
if env_path:
p = Path(env_path)
if p.exists():
return p
current_file = Path(__file__).resolve()
candidate = current_file.parent.parent.parent.parent
if (candidate / "matrices").exists():
return candidate
return current_file.parent.parent.parent
def _ensure_dirs(self):
"""确保必要的目录存在"""
self.novel_rag_dir.mkdir(parents=True, exist_ok=True)
if self.ws_rag_dir:
self.ws_rag_dir.mkdir(parents=True, exist_ok=True)
if self.ws_shared_dir:
self.ws_shared_dir.mkdir(parents=True, exist_ok=True)
if self.ws_catharsis_dir:
self.ws_catharsis_dir.mkdir(parents=True, exist_ok=True)
if self.ws_genres_dir:
self.ws_genres_dir.mkdir(parents=True, exist_ok=True)
if self.ws_learned_dir:
self.ws_learned_dir.mkdir(parents=True, exist_ok=True)
if self.proj_rag_dir:
self.proj_rag_dir.mkdir(parents=True, exist_ok=True)
if self.proj_shared_dir:
self.proj_shared_dir.mkdir(parents=True, exist_ok=True)
if self.proj_catharsis_dir:
self.proj_catharsis_dir.mkdir(parents=True, exist_ok=True)
if self.proj_genres_dir:
self.proj_genres_dir.mkdir(parents=True, exist_ok=True)
if self.proj_learned_dir:
self.proj_learned_dir.mkdir(parents=True, exist_ok=True)
def _load_config(self):
"""加载配置"""
if self.config is not None:
self._extract_embed_config(self.config)
return
try:
from data_modules.config import DataModulesConfig
self.config = DataModulesConfig.from_project_root(self.project_root)
self._extract_embed_config(self.config)
return
except Exception:
pass
self._embed_base_url = os.getenv("EMBED_BASE_URL", "")
self._embed_model = os.getenv("EMBED_MODEL", "")
self._embed_api_key = os.getenv("EMBED_API_KEY", "")
def _extract_embed_config(self, config):
"""从配置对象提取 embedding 配置"""
self._embed_base_url = getattr(config, 'embed_base_url', "") or os.getenv("EMBED_BASE_URL", "")
self._embed_model = getattr(config, 'embed_model', "") or os.getenv("EMBED_MODEL", "")
self._embed_api_key = getattr(config, 'embed_api_key', "") or os.getenv("EMBED_API_KEY", "")
async def search(
self,
query: str,
top_k: int = 5,
layers: Optional[List[RAGLayer]] = None,
chunk_type: Optional[str] = None,
) -> List[CrossProjectSearchResult]:
"""四层检索(向下继承)"""
if layers is None:
layers = [RAGLayer.NOVEL, RAGLayer.WORKSPACE, RAGLayer.PROJECT, RAGLayer.PLUGIN]
all_results = []
tasks_with_layers = []
if RAGLayer.NOVEL in layers:
tasks_with_layers.append((RAGLayer.NOVEL, self._search_novel(query, top_k, chunk_type)))
if RAGLayer.WORKSPACE in layers and self.ws_rag_dir:
tasks_with_layers.append((RAGLayer.WORKSPACE, self._search_workspace(query, top_k, chunk_type)))
if RAGLayer.PROJECT in layers and self.proj_rag_dir:
tasks_with_layers.append((RAGLayer.PROJECT, self._search_project(query, top_k, chunk_type)))
if RAGLayer.PLUGIN in layers:
tasks_with_layers.append((RAGLayer.PLUGIN, self._search_plugin(query, top_k, chunk_type)))
if tasks_with_layers:
tasks = [t[1] for t in tasks_with_layers]
layer_results = await asyncio.gather(*tasks)
for (layer, _), results in zip(tasks_with_layers, layer_results):
if results:
for r in results:
r.score *= self.weights.get(layer, 1.0)
all_results.append(r)
all_results.sort(key=lambda x: x.score, reverse=True)
return all_results[:top_k]
async def _search_novel(self, query: str, top_k: int, chunk_type: Optional[str]) -> List[CrossProjectSearchResult]:
"""检索小说私有 RAG"""
if not self.novel_vectors_db.exists():
return []
if self._embed_api_key:
return await self._novel_vector_search(query, top_k, chunk_type)
return await asyncio.to_thread(self._novel_keyword_search, query, top_k, chunk_type)
def _novel_keyword_search(self, query: str, top_k: int, chunk_type: Optional[str]) -> List[CrossProjectSearchResult]:
"""小说关键词检索"""
try:
conn = sqlite3.connect(str(self.novel_vectors_db))
cursor = conn.cursor()
if chunk_type:
cursor.execute("""
SELECT chunk_id, chapter, content, chunk_type, source_file
FROM vectors WHERE chunk_type = ? AND content LIKE ?
ORDER BY chapter DESC LIMIT ?
""", (chunk_type, f"%{query}%", top_k))
else:
cursor.execute("""
SELECT chunk_id, chapter, content, chunk_type, source_file
FROM vectors WHERE content LIKE ?
ORDER BY chapter DESC LIMIT ?
""", (f"%{query}%", top_k))
rows = cursor.fetchall()
conn.close()
keywords = self._extract_keywords(query)
results = []
for row in rows:
content = row[2] or ""
matches = sum(1 for kw in keywords if kw in content)
score = matches / max(len(keywords), 1) * 100
results.append(CrossProjectSearchResult(
chunk_id=row[0], content=content[:500], score=score,
source_layer=RAGLayer.NOVEL, source_project=self.novel_root.name,
chapter=row[1], chunk_type=row[3], metadata={"source_file": row[4]}
))
results.sort(key=lambda x: x.score, reverse=True)
return results[:top_k]
except Exception:
return []
async def _novel_vector_search(self, query: str, top_k: int, chunk_type: Optional[str]) -> List[CrossProjectSearchResult]:
"""小说向量检索"""
try:
embeddings = await self._embed_texts([query])
if not embeddings:
return await asyncio.to_thread(self._novel_keyword_search, query, top_k, chunk_type)
query_embedding = embeddings[0]
conn = sqlite3.connect(str(self.novel_vectors_db))
cursor = conn.cursor()
if chunk_type:
cursor.execute("""
SELECT chunk_id, chapter, content, embedding, chunk_type, source_file
FROM vectors WHERE chunk_type = ?
""", (chunk_type,))
else:
cursor.execute("SELECT chunk_id, chapter, content, embedding, chunk_type, source_file FROM vectors")
rows = cursor.fetchall()
conn.close()
results = []
for row in rows:
if not row[3]:
continue
embedding = self._deserialize_embedding(row[3])
score = self._cosine_similarity(query_embedding, embedding)
results.append(CrossProjectSearchResult(
chunk_id=row[0], content=row[2][:500] if row[2] else "", score=score,
source_layer=RAGLayer.NOVEL, source_project=self.novel_root.name,
chapter=row[1], chunk_type=row[4], metadata={"source_file": row[5]}
))
results.sort(key=lambda x: x.score, reverse=True)
return results[:top_k]
except Exception:
return await asyncio.to_thread(self._novel_keyword_search, query, top_k, chunk_type)
async def _search_workspace(self, query: str, top_k: int, chunk_type: Optional[str]) -> List[CrossProjectSearchResult]:
"""检索工作空间共享 RAG"""
results = []
# 检索 catharsis 模板
if self.ws_catharsis_dir and self.ws_catharsis_dir.exists():
for model_file in self.ws_catharsis_dir.glob("*.md"):
try:
content = model_file.read_text(encoding="utf-8")
keywords = self._extract_keywords(query)
matches = sum(1 for kw in keywords if kw in content[:1000])
if matches > 0:
score = matches / len(keywords) * 70
results.append(CrossProjectSearchResult(
chunk_id=f"workspace:{model_file.stem}", content=content[:500], score=score,
source_layer=RAGLayer.WORKSPACE, source_project="workspace",
chapter=None, chunk_type="catharsis_model"
))
except Exception:
continue
# 检索题材库
if self.ws_genres_dir and self.ws_genres_dir.exists():
keywords = self._extract_keywords(query)
for genre_file in self.ws_genres_dir.glob("**/*.md"):
try:
content = genre_file.read_text(encoding="utf-8")
matches = sum(1 for kw in keywords if kw in content[:1000])
if matches > 0:
score = matches / len(keywords) * 50
results.append(CrossProjectSearchResult(
chunk_id=f"workspace_genre:{genre_file.stem}", content=content[:500], score=score,
source_layer=RAGLayer.WORKSPACE, source_project="workspace",
chapter=None, chunk_type="genre_template"
))
except Exception:
continue
# 检索学习成果
if self.ws_learned_dir and self.ws_learned_dir.exists():
results.extend(await self._search_learned_dir(self.ws_learned_dir, query, top_k, RAGLayer.WORKSPACE))
return results[:top_k]
async def _search_project(self, query: str, top_k: int, chunk_type: Optional[str]) -> List[CrossProjectSearchResult]:
"""检索工程共享 RAG"""
results = []
if self.proj_catharsis_dir and self.proj_catharsis_dir.exists():
for model_file in self.proj_catharsis_dir.glob("*.md"):
try:
content = model_file.read_text(encoding="utf-8")
keywords = self._extract_keywords(query)
matches = sum(1 for kw in keywords if kw in content[:1000])
if matches > 0:
score = matches / len(keywords) * 50
results.append(CrossProjectSearchResult(
chunk_id=f"project:{model_file.stem}", content=content[:500], score=score,
source_layer=RAGLayer.PROJECT, source_project="project",
chapter=None, chunk_type="catharsis_model"
))
except Exception:
continue
if self.proj_genres_dir and self.proj_genres_dir.exists():
keywords = self._extract_keywords(query)
for genre_file in self.proj_genres_dir.glob("**/*.md"):
try:
content = genre_file.read_text(encoding="utf-8")
matches = sum(1 for kw in keywords if kw in content[:1000])
if matches > 0:
score = matches / len(keywords) * 30
results.append(CrossProjectSearchResult(
chunk_id=f"project_genre:{genre_file.stem}", content=content[:500], score=score,
source_layer=RAGLayer.PROJECT, source_project="project",
chapter=None, chunk_type="genre_template"
))
except Exception:
continue
if self.proj_learned_dir and self.proj_learned_dir.exists():
results.extend(await self._search_learned_dir(self.proj_learned_dir, query, top_k, RAGLayer.PROJECT))
return results[:top_k]
async def _search_plugin(self, query: str, top_k: int, chunk_type: Optional[str]) -> List[CrossProjectSearchResult]:
"""检索插件内置 RAG"""
results = []
if self.plugin_catharsis_dir and self.plugin_catharsis_dir.exists():
for model_file in self.plugin_catharsis_dir.glob("*.md"):
try:
content = model_file.read_text(encoding="utf-8")
keywords = self._extract_keywords(query)
matches = sum(1 for kw in keywords if kw in content[:1000])
if matches > 0:
score = matches / len(keywords) * 20
results.append(CrossProjectSearchResult(
chunk_id=f"plugin:{model_file.stem}", content=content[:500], score=score,
source_layer=RAGLayer.PLUGIN, source_project="plugin",
chapter=None, chunk_type="catharsis_model"
))
except Exception:
continue
return results[:top_k]
async def _search_learned_dir(self, learned_dir: Path, query: str, top_k: int, layer: RAGLayer) -> List[CrossProjectSearchResult]:
"""检索学习成果目录"""
results = []
keywords = self._extract_keywords(query)
for pattern_file in learned_dir.glob("**/*.json"):
try:
data = json.loads(pattern_file.read_text(encoding="utf-8"))
content = json.dumps(data, ensure_ascii=False)
matches = sum(1 for kw in keywords if kw in content)
if matches > 0:
score = matches / max(len(keywords), 1) * 50
results.append(CrossProjectSearchResult(
chunk_id=f"learned:{pattern_file.stem}", content=content[:500], score=score,
source_layer=layer, source_project=data.get("source_project", "unknown"),
chapter=data.get("source_chapter"), chunk_type="learned_pattern", metadata=data
))
except Exception:
continue
return results[:top_k]
def _extract_keywords(self, query: str) -> List[str]:
"""提取关键词"""
import re
chinese = re.findall(r'[\u4e00-\u9fff]{2,8}', query)
english = re.findall(r'[a-zA-Z]{2,}', query.lower())
return chinese + english
def _cosine_similarity(self, a: List[float], b: List[float]) -> float:
dot_product = sum(x * y for x, y in zip(a, b))
norm_a = sum(x * x for x in a) ** 0.5
norm_b = sum(x * x for x in b) ** 0.5
if norm_a == 0 or norm_b == 0:
return 0.0
return dot_product / (norm_a * norm_b)
def _deserialize_embedding(self, data: bytes) -> List[float]:
import struct
count = len(data) // 4
return list(struct.unpack(f"{count}f", data))
async def _embed_texts(self, texts: List[str]) -> Optional[List[List[float]]]:
"""调用 Embedding API"""
if not self._embed_api_key or not texts:
return None
try:
import aiohttp
headers = {"Authorization": f"Bearer {self._embed_api_key}", "Content-Type": "application/json"}
payload = {"model": self._embed_model, "input": texts}
async with aiohttp.ClientSession() as session:
async with session.post(
f"{self._embed_base_url}/embeddings", headers=headers, json=payload,
timeout=aiohttp.ClientTimeout(total=60)
) as resp:
if resp.status == 200:
result = await resp.json()
return [item["embedding"] for item in result["data"]]
except Exception:
pass
return None
# ==================== 存储接口 ====================
def store_learned_pattern(self, pattern: LearnedPattern, layer: RAGLayer = RAGLayer.NOVEL) -> bool:
"""存储学习到的模式"""
if layer == RAGLayer.NOVEL:
return self._store_novel_learned(pattern)
elif layer == RAGLayer.WORKSPACE:
return self._store_workspace_learned(pattern)
elif layer == RAGLayer.PROJECT:
return self._store_project_learned(pattern)
return False
def _store_novel_learned(self, pattern: LearnedPattern) -> bool:
"""存储到小说私有库"""
self._init_learned_db(self.novel_learned_db)
try:
conn = sqlite3.connect(str(self.novel_learned_db))
cursor = conn.cursor()
cursor.execute("""
INSERT OR REPLACE INTO learned_patterns
(pattern_id, pattern_type, title, description, tension_curve, catharsis_model,
structure, hot_spots, style_tags, source_project, source_chapter, learned_at, usage_count, metadata)
VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)
""", (
pattern.pattern_id, pattern.pattern_type, pattern.title, pattern.description,
json.dumps(pattern.tension_curve), pattern.catharsis_model, json.dumps(pattern.structure),
json.dumps(pattern.hot_spots), json.dumps(pattern.style_tags), pattern.source_project,
pattern.source_chapter, pattern.learned_at, pattern.usage_count, json.dumps(pattern.metadata)
))
conn.commit()
conn.close()
return True
except Exception:
return False
def _store_workspace_learned(self, pattern: LearnedPattern) -> bool:
"""存储到工作空间共享"""
if not self.ws_learned_dir:
return False
pattern_file = self.ws_learned_dir / f"{pattern.pattern_id}.json"
return self._write_pattern_file(pattern_file, pattern)
def _store_project_learned(self, pattern: LearnedPattern) -> bool:
"""存储到工程共享"""
if not self.proj_learned_dir:
return False
pattern_file = self.proj_learned_dir / f"{pattern.pattern_id}.json"
return self._write_pattern_file(pattern_file, pattern)
def _write_pattern_file(self, path: Path, pattern: LearnedPattern) -> bool:
try:
path.parent.mkdir(parents=True, exist_ok=True)
data = {
"pattern_id": pattern.pattern_id, "pattern_type": pattern.pattern_type,
"title": pattern.title, "description": pattern.description,
"tension_curve": pattern.tension_curve, "catharsis_model": pattern.catharsis_model,
"structure": pattern.structure, "hot_spots": pattern.hot_spots,
"style_tags": pattern.style_tags, "source_project": pattern.source_project,
"source_chapter": pattern.source_chapter, "learned_at": pattern.learned_at,
"usage_count": pattern.usage_count, "metadata": pattern.metadata
}
path.write_text(json.dumps(data, ensure_ascii=False, indent=2), encoding="utf-8")
return True
except Exception:
return False
def _init_learned_db(self, db_path: Path):
"""初始化学习库"""
if db_path.exists():
return
conn = sqlite3.connect(str(db_path))
cursor = conn.cursor()
cursor.execute("""
CREATE TABLE IF NOT EXISTS learned_patterns (
pattern_id TEXT PRIMARY KEY, pattern_type TEXT NOT NULL, title TEXT NOT NULL,
description TEXT, tension_curve TEXT, catharsis_model TEXT, structure TEXT,
hot_spots TEXT, style_tags TEXT, source_project TEXT, source_chapter INTEGER,
learned_at TEXT, usage_count INTEGER DEFAULT 0, metadata TEXT
)
""")
conn.commit()
conn.close()
# ==================== 项目索引 ====================
@staticmethod
def get_projects_index(system_root: Path) -> Dict[str, Any]:
projects_file = system_root / "projects.json"
if projects_file.exists():
return json.loads(projects_file.read_text(encoding="utf-8"))
return {"projects": []}
@staticmethod
def register_project(system_root: Path, project_path: Path, project_info: Dict[str, Any]) -> bool:
system_root = Path(system_root)
system_root.mkdir(parents=True, exist_ok=True)
projects_file = system_root / "projects.json"
data = CrossProjectRAG.get_projects_index(system_root)
project_path_str = str(project_path.resolve())
projects = data.get("projects", [])
for i, p in enumerate(projects):
if p.get("path") == project_path_str:
projects[i] = project_info
break
else:
projects.append(project_info)
data["projects"] = projects
data["last_updated"] = datetime.now().isoformat()
try:
projects_file.write_text(json.dumps(data, ensure_ascii=False, indent=2), encoding="utf-8")
return True
except Exception:
return False
if __name__ == "__main__":
import argparse
import sys
if sys.platform == "win32":
enable_windows_utf8_stdio()
parser = argparse.ArgumentParser(description="Cross-Project RAG CLI")
parser.add_argument("--project-root", type=str, required=True)
parser.add_argument("--workspace-root", type=str)
parser.add_argument("--project-root-dir", type=str)
subparsers = parser.add_subparsers(dest="command")
search_parser = subparsers.add_parser("search")
search_parser.add_argument("--query", required=True)
search_parser.add_argument("--top-k", type=int, default=5)
search_parser.add_argument("--layers", type=str, default="novel,workspace,project,plugin")
args = parser.parse_args()
if not args.project_root:
print("Error: --project-root is required")
sys.exit(1)
rag = CrossProjectRAG(
project_root=Path(args.project_root).resolve(),
workspace_root=Path(args.workspace_root).resolve() if args.workspace_root else None,
project_root_dir=Path(args.project_root_dir).resolve() if args.project_root_dir else None,
)
if args.command == "search":
layer_map = {"novel": RAGLayer.NOVEL, "workspace": RAGLayer.WORKSPACE,
"project": RAGLayer.PROJECT, "plugin": RAGLayer.PLUGIN}
layers = [layer_map[l.strip()] for l in args.layers.split(",") if l.strip() in layer_map]
results = asyncio.run(rag.search(args.query, args.top_k, layers))
print(f"\n=== Search Results ({len(results)}) ===")
for r in results:
print(f"\n[{r.source_layer.value}] {r.chunk_id} (score: {r.score:.2f})")
print(f"Source: {r.source_project or 'unknown'}")
if r.chapter:
print(f"Chapter: {r.chapter}")
print(f"Content: {r.content[:200]}...")
else:
parser.print_help()
+274
View File
@@ -0,0 +1,274 @@
#!/usr/bin/env python3
# -*- coding: utf-8 -*-
"""
Entity Linker - 实体消歧辅助模块 (v5.4)
为 Data Agent 提供实体消歧的辅助功能:
- 置信度判断
- 别名索引管理 (通过 index.db aliases 表)
- 消歧结果记录
v5.1 变更(v5.4 沿用):
- 别名存储从 state.json 迁移到 index.db aliases 表
- 使用 IndexManager 进行别名读写
- 移除对 state.json 的直接操作
"""
from typing import Dict, List, Optional, Tuple
from dataclasses import dataclass, field
from .config import get_config
from .index_manager import IndexManager
from .observability import safe_log_tool_call
@dataclass
class DisambiguationResult:
"""消歧结果"""
mention: str
entity_id: Optional[str]
confidence: float
candidates: List[str] = field(default_factory=list)
adopted: bool = False
warning: Optional[str] = None
class EntityLinker:
"""实体链接器 - 辅助 Data Agent 进行实体消歧 (v5.1 SQLite,v5.4 沿用)"""
def __init__(self, config=None):
self.config = config or get_config()
self._index_manager = IndexManager(self.config)
# ==================== 别名管理 (v5.1 SQLite,v5.4 沿用) ====================
def register_alias(self, entity_id: str, alias: str, entity_type: str = "角色") -> bool:
"""注册新别名(v5.1 引入:写入 index.db aliases 表)"""
if not alias or not entity_id:
return False
return self._index_manager.register_alias(alias, entity_id, entity_type)
def lookup_alias(self, mention: str, entity_type: str = None) -> Optional[str]:
"""查找别名对应的实体ID(返回第一个匹配,可选按类型过滤)"""
entries = self._index_manager.get_entities_by_alias(mention)
if not entries:
return None
if entity_type:
for entry in entries:
if entry.get("type") == entity_type:
return entry.get("id")
return None
else:
return entries[0].get("id") if entries else None
def lookup_alias_all(self, mention: str) -> List[Dict]:
"""查找别名对应的所有实体(一对多)"""
entries = self._index_manager.get_entities_by_alias(mention)
return [{"type": e.get("type"), "id": e.get("id")} for e in entries]
def get_all_aliases(self, entity_id: str, entity_type: str = None) -> List[str]:
"""获取实体的所有别名"""
return self._index_manager.get_entity_aliases(entity_id)
# ==================== 置信度判断 ====================
def evaluate_confidence(self, confidence: float) -> Tuple[str, bool, Optional[str]]:
"""
评估置信度,返回 (action, adopt, warning)
- action: "auto" | "warn" | "manual"
- adopt: 是否采用
- warning: 警告信息
"""
if confidence >= self.config.extraction_confidence_high:
return ("auto", True, None)
elif confidence >= self.config.extraction_confidence_medium:
return ("warn", True, f"中置信度匹配 (confidence: {confidence:.2f})")
else:
return ("manual", False, f"需人工确认 (confidence: {confidence:.2f})")
def process_uncertain(
self,
mention: str,
candidates: List[str],
suggested: str,
confidence: float,
context: str = ""
) -> DisambiguationResult:
"""
处理不确定的实体匹配
返回消歧结果,包含是否采用、警告信息等
"""
action, adopt, warning = self.evaluate_confidence(confidence)
result = DisambiguationResult(
mention=mention,
entity_id=suggested if adopt else None,
confidence=confidence,
candidates=candidates,
adopted=adopt,
warning=warning
)
return result
# ==================== 批量处理 ====================
def process_extraction_result(
self,
uncertain_items: List[Dict]
) -> Tuple[List[DisambiguationResult], List[str]]:
"""
处理 AI 提取结果中的 uncertain 项
返回 (results, warnings)
"""
results = []
warnings = []
for item in uncertain_items:
result = self.process_uncertain(
mention=item.get("mention", ""),
candidates=item.get("candidates", []),
suggested=item.get("suggested", ""),
confidence=item.get("confidence", 0.0),
context=item.get("context", "")
)
results.append(result)
if result.warning:
warnings.append(f"{result.mention} → {result.entity_id}: {result.warning}")
return results, warnings
def register_new_entities(
self,
new_entities: List[Dict]
) -> List[str]:
"""
注册新实体的别名 (v5.1 引入,v5.4 沿用)
返回注册的实体ID列表
"""
registered = []
for entity in new_entities:
entity_id = entity.get("suggested_id") or entity.get("id")
if not entity_id or entity_id == "NEW":
continue
entity_type = entity.get("type", "角色")
# 注册主名称
name = entity.get("name", "")
if name:
self.register_alias(entity_id, name, entity_type)
# 注册提及方式
for mention in entity.get("mentions", []):
if mention and mention != name:
self.register_alias(entity_id, mention, entity_type)
registered.append(entity_id)
return registered
# ==================== CLI 接口 ====================
def main():
import argparse
import sys
from .cli_output import print_success, print_error
from .cli_args import normalize_global_project_root
from .index_manager import IndexManager
parser = argparse.ArgumentParser(description="Entity Linker CLI (v5.4 SQLite)")
parser.add_argument("--project-root", type=str, help="项目根目录")
subparsers = parser.add_subparsers(dest="command")
# 注册别名
register_parser = subparsers.add_parser("register-alias")
register_parser.add_argument("--entity", required=True, help="实体ID")
register_parser.add_argument("--alias", required=True, help="别名")
register_parser.add_argument("--type", default="角色", help="实体类型(默认:角色)")
# 查找别名
lookup_parser = subparsers.add_parser("lookup")
lookup_parser.add_argument("--mention", required=True, help="提及文本")
lookup_parser.add_argument("--type", help="按类型过滤")
# 查找所有匹配(一对多)
lookup_all_parser = subparsers.add_parser("lookup-all")
lookup_all_parser.add_argument("--mention", required=True, help="提及文本")
# 列出别名
list_parser = subparsers.add_parser("list-aliases")
list_parser.add_argument("--entity", required=True, help="实体ID")
list_parser.add_argument("--type", help="实体类型")
argv = normalize_global_project_root(sys.argv[1:])
args = parser.parse_args(argv)
# 初始化
config = None
if args.project_root:
# 允许传入“工作区根目录”,统一解析到真正的 book project_root(必须包含 .noma/state.json)
from project_locator import resolve_project_root
from .config import DataModulesConfig
resolved_root = resolve_project_root(args.project_root)
config = DataModulesConfig.from_project_root(resolved_root)
linker = EntityLinker(config)
logger = IndexManager(config)
tool_name = f"entity_linker:{args.command or 'unknown'}"
def emit_success(data=None, message: str = "ok"):
print_success(data, message=message)
safe_log_tool_call(logger, tool_name=tool_name, success=True)
def emit_error(code: str, message: str, suggestion: str | None = None):
print_error(code, message, suggestion=suggestion)
safe_log_tool_call(
logger,
tool_name=tool_name,
success=False,
error_code=code,
error_message=message,
)
if args.command == "register-alias":
entity_type = getattr(args, "type", "角色")
success = linker.register_alias(args.entity, args.alias, entity_type)
if success:
emit_success({"entity": args.entity, "alias": args.alias, "type": entity_type}, message="alias_registered")
else:
emit_error("ALIAS_EXISTS", "注册失败或已存在")
elif args.command == "lookup":
entity_type = getattr(args, "type", None)
entity_id = linker.lookup_alias(args.mention, entity_type)
if entity_id:
emit_success({"mention": args.mention, "entity": entity_id}, message="lookup")
else:
emit_error("NOT_FOUND", f"未找到别名: {args.mention}")
elif args.command == "lookup-all":
matches = linker.lookup_alias_all(args.mention)
emit_success(matches, message="lookup_all")
elif args.command == "list-aliases":
entity_type = getattr(args, "type", None)
aliases = linker.get_all_aliases(args.entity, entity_type)
emit_success(aliases, message="aliases")
else:
emit_error("UNKNOWN_COMMAND", "未指定有效命令", suggestion="请查看 --help")
if __name__ == "__main__":
main()
@@ -0,0 +1,66 @@
#!/usr/bin/env python3
# -*- coding: utf-8 -*-
"""
Genre alias normalization and profile key mapping.
"""
from __future__ import annotations
GENRE_INPUT_ALIASES: dict[str, str] = {
"修仙/玄幻": "修仙",
"玄幻修仙": "修仙",
"玄幻": "修仙",
"修真": "修仙",
"都市修真": "都市异能",
"都市高武": "高武",
"都市奇闻": "都市脑洞",
"古言脑洞": "古言",
"游戏电竞": "电竞",
"电竞文": "电竞",
"直播": "直播文",
"直播带货": "直播文",
"主播": "直播文",
"克系": "克苏鲁",
"克系悬疑": "克苏鲁",
}
GENRE_PROFILE_KEY_ALIASES: dict[str, str] = {
"修仙": "xianxia",
"修仙/玄幻": "xianxia",
"玄幻": "xianxia",
"爽文/系统流": "shuangwen",
"高武": "xianxia",
"西幻": "xianxia",
"都市异能": "urban-power",
"都市脑洞": "urban-power",
"都市日常": "urban-power",
"狗血言情": "romance",
"古言": "romance",
"青春甜宠": "romance",
"替身文": "substitute",
"规则怪谈": "rules-mystery",
"悬疑脑洞": "mystery",
"悬疑灵异": "mystery",
"知乎短篇": "zhihu-short",
"电竞": "esports",
"直播文": "livestream",
"克苏鲁": "cosmic-horror",
}
def normalize_genre_token(token: str) -> str:
value = str(token or "").strip()
if not value:
return ""
return GENRE_INPUT_ALIASES.get(value, value)
def to_profile_key(genre: str) -> str:
value = str(genre or "").strip()
if not value:
return ""
normalized = normalize_genre_token(value)
return GENRE_PROFILE_KEY_ALIASES.get(normalized, normalized.lower())
@@ -0,0 +1,107 @@
#!/usr/bin/env python3
# -*- coding: utf-8 -*-
"""
Genre profile parsing helpers for ContextManager.
"""
from __future__ import annotations
import re
from typing import List
from .genre_aliases import normalize_genre_token
def parse_genre_tokens(
genre_raw: str,
*,
support_composite: bool,
separators: tuple[str, ...],
) -> List[str]:
text = str(genre_raw or "").strip()
if not text:
return []
if not support_composite:
normalized_single = normalize_genre_token(text)
return [normalized_single] if normalized_single else [text]
pattern = "|".join(re.escape(str(token)) for token in separators if str(token))
if not pattern:
normalized_single = normalize_genre_token(text)
return [normalized_single] if normalized_single else [text]
tokens = [chunk.strip() for chunk in re.split(pattern, text) if chunk and chunk.strip()]
deduped: List[str] = []
seen = set()
for token in tokens:
normalized_token = normalize_genre_token(token)
if not normalized_token:
continue
lower = normalized_token.lower()
if lower in seen:
continue
seen.add(lower)
deduped.append(normalized_token)
if deduped:
return deduped
fallback_token = normalize_genre_token(text)
return [fallback_token] if fallback_token else [text]
def extract_genre_section(text: str, genre: str) -> str:
if not text:
return ""
lines = text.splitlines()
capture: List[str] = []
active = False
target = genre.strip().lower()
for line in lines:
normalized = line.strip().lower()
if normalized.startswith("## ") or normalized.startswith("### "):
if active:
break
active = target in normalized
if active:
capture.append(line)
continue
if active:
capture.append(line)
if capture:
return "\n".join(capture).strip()
return "\n".join(lines[:80]).strip()
def extract_markdown_refs(text: str, max_items: int = 8) -> List[str]:
if not text:
return []
refs: List[str] = []
for line in text.splitlines():
row = line.strip().lstrip("-*").strip()
if not row or row.startswith("#"):
continue
refs.append(row)
if len(refs) >= max(1, max_items):
break
return refs
def build_composite_genre_hints(genres: List[str], refs: List[str]) -> List[str]:
if len(genres) <= 1:
return []
primary = genres[0]
secondaries = genres[1:]
hints: List[str] = []
hints.append(
f"以“{primary}”作为主引擎推进主线,每章至少保留1处“{'/'.join(secondaries)}”特征表达。"
)
if refs:
hints.append(f"复合题材执行参考:{refs[0]}")
hints.append("主辅题材冲突时,优先保证主题材读者承诺,辅题材用于制造新鲜感。")
return hints
@@ -0,0 +1,302 @@
#!/usr/bin/env python3
# -*- coding: utf-8 -*-
"""
IndexChapterMixin extracted from IndexManager.
"""
from __future__ import annotations
import json
from datetime import datetime
from typing import Any, Dict, List, Optional
class IndexChapterMixin:
def add_chapter(self, meta: ChapterMeta):
"""添加/更新章节元数据"""
with self._get_conn() as conn:
cursor = conn.cursor()
cursor.execute(
"""
INSERT OR REPLACE INTO chapters
(chapter, title, location, word_count, characters, summary)
VALUES (?, ?, ?, ?, ?, ?)
""",
(
meta.chapter,
meta.title,
meta.location,
meta.word_count,
json.dumps(meta.characters, ensure_ascii=False),
meta.summary,
),
)
conn.commit()
def get_chapter(self, chapter: int) -> Optional[Dict]:
"""获取章节元数据"""
with self._get_conn() as conn:
cursor = conn.cursor()
cursor.execute("SELECT * FROM chapters WHERE chapter = ?", (chapter,))
row = cursor.fetchone()
if row:
return self._row_to_dict(row, parse_json=["characters"])
return None
def get_recent_chapters(self, limit: int = None) -> List[Dict]:
"""获取最近章节"""
if limit is None:
limit = self.config.query_recent_chapters_limit
with self._get_conn() as conn:
cursor = conn.cursor()
cursor.execute(
"""
SELECT * FROM chapters
ORDER BY chapter DESC
LIMIT ?
""",
(limit,),
)
return [
self._row_to_dict(row, parse_json=["characters"])
for row in cursor.fetchall()
]
# ==================== 场景操作 ====================
def add_scenes(self, chapter: int, scenes: List[SceneMeta]):
"""添加章节场景"""
with self._get_conn() as conn:
cursor = conn.cursor()
# 先删除该章节旧场景
cursor.execute("DELETE FROM scenes WHERE chapter = ?", (chapter,))
# 插入新场景
for scene in scenes:
cursor.execute(
"""
INSERT INTO scenes
(chapter, scene_index, start_line, end_line, location, summary, characters)
VALUES (?, ?, ?, ?, ?, ?, ?)
""",
(
scene.chapter,
scene.scene_index,
scene.start_line,
scene.end_line,
scene.location,
scene.summary,
json.dumps(scene.characters, ensure_ascii=False),
),
)
conn.commit()
def get_scenes(self, chapter: int) -> List[Dict]:
"""获取章节场景"""
with self._get_conn() as conn:
cursor = conn.cursor()
cursor.execute(
"""
SELECT * FROM scenes
WHERE chapter = ?
ORDER BY scene_index
""",
(chapter,),
)
return [
self._row_to_dict(row, parse_json=["characters"])
for row in cursor.fetchall()
]
def search_scenes_by_location(self, location: str, limit: int = None) -> List[Dict]:
"""按地点搜索场景"""
if limit is None:
limit = self.config.query_scenes_by_location_limit
with self._get_conn() as conn:
cursor = conn.cursor()
cursor.execute(
"""
SELECT * FROM scenes
WHERE location LIKE ?
ORDER BY chapter DESC
LIMIT ?
""",
(f"%{location}%", limit),
)
return [
self._row_to_dict(row, parse_json=["characters"])
for row in cursor.fetchall()
]
# ==================== 出场记录操作 ====================
def record_appearance(
self,
entity_id: str,
chapter: int,
mentions: List[str],
confidence: float = 1.0,
skip_if_exists: bool = False,
):
"""记录实体出场
Args:
entity_id: 实体ID
chapter: 章节号
mentions: 提及列表
confidence: 置信度
skip_if_exists: 如果为True,当记录已存在时跳过(避免覆盖已有mentions)
"""
with self._get_conn() as conn:
cursor = conn.cursor()
if skip_if_exists:
# 先检查是否已存在
cursor.execute(
"SELECT 1 FROM appearances WHERE entity_id = ? AND chapter = ?",
(entity_id, chapter),
)
if cursor.fetchone():
return # 已存在,跳过
cursor.execute(
"""
INSERT OR REPLACE INTO appearances
(entity_id, chapter, mentions, confidence)
VALUES (?, ?, ?, ?)
""",
(
entity_id,
chapter,
json.dumps(mentions, ensure_ascii=False),
confidence,
),
)
conn.commit()
def get_entity_appearances(self, entity_id: str, limit: int = None) -> List[Dict]:
"""获取实体出场记录"""
if limit is None:
limit = self.config.query_entity_appearances_limit
with self._get_conn() as conn:
cursor = conn.cursor()
cursor.execute(
"""
SELECT * FROM appearances
WHERE entity_id = ?
ORDER BY chapter DESC
LIMIT ?
""",
(entity_id, limit),
)
return [
self._row_to_dict(row, parse_json=["mentions"])
for row in cursor.fetchall()
]
def get_recent_appearances(self, limit: int = None) -> List[Dict]:
"""获取最近出场的实体"""
if limit is None:
limit = self.config.query_recent_appearances_limit
with self._get_conn() as conn:
cursor = conn.cursor()
cursor.execute(
"""
SELECT entity_id, MAX(chapter) as last_chapter, COUNT(*) as total
FROM appearances
GROUP BY entity_id
ORDER BY last_chapter DESC
LIMIT ?
""",
(limit,),
)
return [dict(row) for row in cursor.fetchall()]
def get_chapter_appearances(self, chapter: int) -> List[Dict]:
"""获取某章所有出场实体"""
with self._get_conn() as conn:
cursor = conn.cursor()
cursor.execute(
"""
SELECT * FROM appearances
WHERE chapter = ?
ORDER BY confidence DESC
""",
(chapter,),
)
return [
self._row_to_dict(row, parse_json=["mentions"])
for row in cursor.fetchall()
]
# ==================== v5.1 实体操作 ====================
def process_chapter_data(
self,
chapter: int,
title: str,
location: str,
word_count: int,
entities: List[Dict],
scenes: List[Dict],
) -> Dict[str, int]:
"""
处理章节数据,批量写入索引
返回写入统计
"""
from .index_manager import ChapterMeta, SceneMeta
stats = {"chapters": 0, "scenes": 0, "appearances": 0}
# 提取出场角色
characters = [e.get("id") for e in entities if e.get("type") == "角色"]
# 写入章节元数据
self.add_chapter(
ChapterMeta(
chapter=chapter,
title=title,
location=location,
word_count=word_count,
characters=characters,
summary="", # 可后续由 Data Agent 生成
)
)
stats["chapters"] = 1
# 写入场景
scene_metas = []
for s in scenes:
scene_metas.append(
SceneMeta(
chapter=chapter,
scene_index=s.get("index", 0),
start_line=s.get("start_line", 0),
end_line=s.get("end_line", 0),
location=s.get("location", ""),
summary=s.get("summary", ""),
characters=s.get("characters", []),
)
)
self.add_scenes(chapter, scene_metas)
stats["scenes"] = len(scene_metas)
# 写入出场记录
for entity in entities:
entity_id = entity.get("id")
if entity_id and entity_id != "NEW":
self.record_appearance(
entity_id=entity_id,
chapter=chapter,
mentions=entity.get("mentions", []),
confidence=entity.get("confidence", 1.0),
)
stats["appearances"] += 1
return stats
# ==================== 辅助方法 ====================
@@ -0,0 +1,504 @@
#!/usr/bin/env python3
# -*- coding: utf-8 -*-
"""
IndexDebtMixin extracted from IndexManager.
"""
from __future__ import annotations
import json
from datetime import datetime
from typing import Any, Dict, List, Optional
class IndexDebtMixin:
def create_override_contract(self, contract: OverrideContractMeta) -> int:
"""
创建或更新 Override Contract
使用 SQLite 的 INSERT ... ON CONFLICT ... DO UPDATE 实现原子 UPSERT:
- 并发安全,无需显式锁
- 保持 id 不变,避免 chase_debt.override_contract_id 悬挂
- 完全冻结终态:已 fulfilled/cancelled 的合约所有字段都不会被修改
兼容性:支持 SQLite 3.24+(ON CONFLICT 语法),不依赖 RETURNING(3.35+)
返回合约 ID
"""
with self._get_conn() as conn:
cursor = conn.cursor()
# 使用 ON CONFLICT 实现原子 UPSERT(SQLite 3.24+)
# 终态完全冻结:fulfilled/cancelled 状态下所有字段都保持不变
cursor.execute(
"""
INSERT INTO override_contracts
(chapter, constraint_type, constraint_id, rationale_type,
rationale_text, payback_plan, due_chapter, status)
VALUES (?, ?, ?, ?, ?, ?, ?, ?)
ON CONFLICT(chapter, constraint_type, constraint_id) DO UPDATE SET
rationale_type = CASE
WHEN override_contracts.status IN ('fulfilled', 'cancelled')
THEN override_contracts.rationale_type
ELSE excluded.rationale_type
END,
rationale_text = CASE
WHEN override_contracts.status IN ('fulfilled', 'cancelled')
THEN override_contracts.rationale_text
ELSE excluded.rationale_text
END,
payback_plan = CASE
WHEN override_contracts.status IN ('fulfilled', 'cancelled')
THEN override_contracts.payback_plan
ELSE excluded.payback_plan
END,
due_chapter = CASE
WHEN override_contracts.status IN ('fulfilled', 'cancelled')
THEN override_contracts.due_chapter
ELSE excluded.due_chapter
END,
status = CASE
WHEN override_contracts.status IN ('fulfilled', 'cancelled')
THEN override_contracts.status
ELSE excluded.status
END
""",
(
contract.chapter,
contract.constraint_type,
contract.constraint_id,
contract.rationale_type,
contract.rationale_text,
contract.payback_plan,
contract.due_chapter,
contract.status,
),
)
# 不使用 RETURNING(需要 SQLite 3.35+),改用查询获取 id
cursor.execute(
"""
SELECT id FROM override_contracts
WHERE chapter = ? AND constraint_type = ? AND constraint_id = ?
""",
(contract.chapter, contract.constraint_type, contract.constraint_id),
)
row = cursor.fetchone()
if not row:
# UPSERT 后查不到记录是异常情况,不应发生
raise RuntimeError(
f"Override Contract UPSERT 后无法获取 id: "
f"chapter={contract.chapter}, type={contract.constraint_type}, "
f"id={contract.constraint_id}"
)
contract_id = row[0]
conn.commit()
return contract_id
def get_pending_overrides(self, before_chapter: int = None) -> List[Dict]:
"""获取待偿还的Override Contracts"""
with self._get_conn() as conn:
cursor = conn.cursor()
if before_chapter:
cursor.execute(
"""
SELECT * FROM override_contracts
WHERE status = 'pending' AND due_chapter <= ?
ORDER BY due_chapter ASC
""",
(before_chapter,),
)
else:
cursor.execute("""
SELECT * FROM override_contracts
WHERE status = 'pending'
ORDER BY due_chapter ASC
""")
return [dict(row) for row in cursor.fetchall()]
def get_overdue_overrides(self, current_chapter: int) -> List[Dict]:
"""获取已逾期的Override Contracts"""
with self._get_conn() as conn:
cursor = conn.cursor()
cursor.execute(
"""
SELECT * FROM override_contracts
WHERE status = 'pending' AND due_chapter < ?
ORDER BY due_chapter ASC
""",
(current_chapter,),
)
return [dict(row) for row in cursor.fetchall()]
def fulfill_override(self, contract_id: int) -> bool:
"""标记Override Contract为已偿还"""
with self._get_conn() as conn:
cursor = conn.cursor()
cursor.execute(
"""
UPDATE override_contracts SET
status = 'fulfilled',
fulfilled_at = CURRENT_TIMESTAMP
WHERE id = ?
""",
(contract_id,),
)
conn.commit()
return cursor.rowcount > 0
def get_chapter_overrides(self, chapter: int) -> List[Dict]:
"""获取某章创建的Override Contracts"""
with self._get_conn() as conn:
cursor = conn.cursor()
cursor.execute(
"""
SELECT * FROM override_contracts WHERE chapter = ?
""",
(chapter,),
)
return [dict(row) for row in cursor.fetchall()]
# ==================== v5.3 追读力债务操作 ====================
def create_debt(self, debt: ChaseDebtMeta) -> int:
"""
创建追读力债务
返回债务 ID
"""
with self._get_conn() as conn:
cursor = conn.cursor()
cursor.execute(
"""
INSERT INTO chase_debt
(debt_type, original_amount, current_amount, interest_rate,
source_chapter, due_chapter, override_contract_id, status)
VALUES (?, ?, ?, ?, ?, ?, ?, ?)
""",
(
debt.debt_type,
debt.original_amount,
debt.current_amount,
debt.interest_rate,
debt.source_chapter,
debt.due_chapter,
debt.override_contract_id if debt.override_contract_id else None,
debt.status,
),
)
conn.commit()
debt_id = cursor.lastrowid
# 记录创建事件
self._record_debt_event(
cursor,
debt_id,
"created",
debt.original_amount,
debt.source_chapter,
f"创建债务: {debt.debt_type}",
)
conn.commit()
return debt_id
def get_active_debts(self) -> List[Dict]:
"""获取所有活跃债务"""
with self._get_conn() as conn:
cursor = conn.cursor()
cursor.execute("""
SELECT * FROM chase_debt
WHERE status = 'active'
ORDER BY due_chapter ASC
""")
return [dict(row) for row in cursor.fetchall()]
def get_overdue_debts(self, current_chapter: int) -> List[Dict]:
"""获取已逾期的债务(包括 active 但已过期的,以及已标记为 overdue 的)"""
with self._get_conn() as conn:
cursor = conn.cursor()
cursor.execute(
"""
SELECT * FROM chase_debt
WHERE (status = 'overdue')
OR (status = 'active' AND due_chapter < ?)
ORDER BY due_chapter ASC
""",
(current_chapter,),
)
return [dict(row) for row in cursor.fetchall()]
def get_total_debt_balance(self) -> float:
"""获取总债务余额(包括 active 和 overdue)"""
with self._get_conn() as conn:
cursor = conn.cursor()
cursor.execute("""
SELECT COALESCE(SUM(current_amount), 0) FROM chase_debt
WHERE status IN ('active', 'overdue')
""")
return cursor.fetchone()[0]
def accrue_interest(self, current_chapter: int) -> Dict[str, Any]:
"""
计算利息(每章调用一次)
- 对 active 和 overdue 债务都计息(逾期债务继续累积利息)
- 使用 debt_events 表防止同一章重复计息
- 检查逾期并更新状态
返回: {debts_processed, total_interest, new_overdues, skipped_already_processed}
"""
result = {
"debts_processed": 0,
"total_interest": 0.0,
"new_overdues": 0,
"skipped_already_processed": 0,
}
with self._get_conn() as conn:
cursor = conn.cursor()
# 获取所有未偿还债务(active + overdue 都继续计息)
cursor.execute("""
SELECT * FROM chase_debt WHERE status IN ('active', 'overdue')
""")
debts = cursor.fetchall()
for debt in debts:
debt_id = debt["id"]
current_amount = debt["current_amount"]
interest_rate = debt["interest_rate"]
due_chapter = debt["due_chapter"]
debt_status = debt["status"]
# 检查本章是否已计息(防止重复调用)
cursor.execute(
"""
SELECT 1 FROM debt_events
WHERE debt_id = ? AND chapter = ? AND event_type = 'interest_accrued'
""",
(debt_id, current_chapter),
)
if cursor.fetchone():
result["skipped_already_processed"] += 1
continue
# 计算利息
interest = current_amount * interest_rate
new_amount = current_amount + interest
# 更新债务
cursor.execute(
"""
UPDATE chase_debt SET
current_amount = ?,
updated_at = CURRENT_TIMESTAMP
WHERE id = ?
""",
(new_amount, debt_id),
)
# 记录利息事件
self._record_debt_event(
cursor,
debt_id,
"interest_accrued",
interest,
current_chapter,
f"利息: {interest:.2f} (利率: {interest_rate * 100:.0f}%)",
)
result["debts_processed"] += 1
result["total_interest"] += interest
# 检查是否逾期(仅对 active 状态的债务)
if debt_status == "active" and current_chapter > due_chapter:
cursor.execute(
"""
UPDATE chase_debt SET status = 'overdue'
WHERE id = ? AND status = 'active'
""",
(debt_id,),
)
if cursor.rowcount > 0:
result["new_overdues"] += 1
self._record_debt_event(
cursor,
debt_id,
"overdue",
new_amount,
current_chapter,
f"债务逾期 (截止: 第{due_chapter}章)",
)
conn.commit()
return result
def pay_debt(self, debt_id: int, amount: float, chapter: int) -> Dict[str, Any]:
"""
偿还债务
- 校验 amount > 0
- 完全偿还时,使用原子 UPDATE 检查并标记关联 Override 为 fulfilled
(并发安全:用 NOT EXISTS 子查询确保所有债务都已清零)
返回: {remaining, fully_paid, override_fulfilled}
"""
# 校验偿还金额
if amount <= 0:
return {
"remaining": 0,
"fully_paid": False,
"error": "偿还金额必须大于0",
}
with self._get_conn() as conn:
cursor = conn.cursor()
cursor.execute(
"SELECT current_amount, override_contract_id FROM chase_debt WHERE id = ?",
(debt_id,),
)
row = cursor.fetchone()
if not row:
return {"remaining": 0, "fully_paid": False, "error": "债务不存在"}
current = row["current_amount"]
override_contract_id = row["override_contract_id"]
remaining = max(0, current - amount)
override_fulfilled = False
if remaining == 0:
# 完全偿还
cursor.execute(
"""
UPDATE chase_debt SET
current_amount = 0,
status = 'paid',
updated_at = CURRENT_TIMESTAMP
WHERE id = ?
""",
(debt_id,),
)
self._record_debt_event(
cursor, debt_id, "full_payment", amount, chapter, "债务已完全偿还"
)
# 原子检查并标记 Override 为 fulfilled
# 使用 NOT EXISTS 子查询确保并发安全:只有当确实没有未清债务时才更新
if override_contract_id:
cursor.execute(
"""
UPDATE override_contracts SET
status = 'fulfilled',
fulfilled_at = CURRENT_TIMESTAMP
WHERE id = ?
AND status = 'pending'
AND NOT EXISTS (
SELECT 1 FROM chase_debt
WHERE override_contract_id = ?
AND status IN ('active', 'overdue')
)
""",
(override_contract_id, override_contract_id),
)
if cursor.rowcount > 0:
override_fulfilled = True
else:
# 部分偿还
cursor.execute(
"""
UPDATE chase_debt SET
current_amount = ?,
updated_at = CURRENT_TIMESTAMP
WHERE id = ?
""",
(remaining, debt_id),
)
self._record_debt_event(
cursor,
debt_id,
"partial_payment",
amount,
chapter,
f"部分偿还,剩余: {remaining:.2f}",
)
conn.commit()
return {
"remaining": remaining,
"fully_paid": remaining == 0,
"override_fulfilled": override_fulfilled,
}
def _record_debt_event(
self,
cursor,
debt_id: int,
event_type: str,
amount: float,
chapter: int,
note: str = "",
):
"""记录债务事件(内部方法)"""
cursor.execute(
"""
INSERT INTO debt_events (debt_id, event_type, amount, chapter, note)
VALUES (?, ?, ?, ?, ?)
""",
(debt_id, event_type, amount, chapter, note),
)
def get_debt_history(self, debt_id: int) -> List[Dict]:
"""获取债务的事件历史"""
with self._get_conn() as conn:
cursor = conn.cursor()
cursor.execute(
"""
SELECT * FROM debt_events
WHERE debt_id = ?
ORDER BY created_at ASC
""",
(debt_id,),
)
return [dict(row) for row in cursor.fetchall()]
# ==================== v5.3 章节追读力元数据操作 ====================
def get_debt_summary(self) -> Dict[str, Any]:
"""获取债务汇总信息"""
with self._get_conn() as conn:
cursor = conn.cursor()
# 活跃债务
cursor.execute("""
SELECT COUNT(*) as count, COALESCE(SUM(current_amount), 0) as total
FROM chase_debt WHERE status = 'active'
""")
active = cursor.fetchone()
# 逾期债务
cursor.execute("""
SELECT COUNT(*) as count, COALESCE(SUM(current_amount), 0) as total
FROM chase_debt WHERE status = 'overdue'
""")
overdue = cursor.fetchone()
# 待偿还Override
cursor.execute("""
SELECT COUNT(*) FROM override_contracts WHERE status = 'pending'
""")
pending_overrides = cursor.fetchone()[0]
return {
"active_debts": active["count"],
"active_total": active["total"],
"overdue_debts": overdue["count"],
"overdue_total": overdue["total"],
"pending_overrides": pending_overrides,
"total_balance": active["total"] + overdue["total"],
}
# ==================== 批量操作 ====================
@@ -0,0 +1,985 @@
#!/usr/bin/env python3
# -*- coding: utf-8 -*-
"""
IndexEntityMixin extracted from IndexManager.
"""
from __future__ import annotations
import json
import logging
import re
import sqlite3
from datetime import datetime
from typing import Any, Dict, List, Optional
logger = logging.getLogger(__name__)
class IndexEntityMixin:
def upsert_entity(self, entity: EntityMeta, update_metadata: bool = False) -> bool:
"""
插入或更新实体 (智能合并)
- 新实体: 直接插入
- 已存在: 更新 current_json, last_appearance, updated_at
- update_metadata=True: 同时更新 canonical_name/tier/desc/is_protagonist/is_archived
返回是否为新实体
"""
with self._get_conn() as conn:
cursor = conn.cursor()
# 检查是否存在
cursor.execute(
"SELECT id, current_json FROM entities WHERE id = ?", (entity.id,)
)
existing = cursor.fetchone()
if existing:
# 已存在: 智能合并 current_json
old_current = {}
if existing["current_json"]:
try:
old_current = json.loads(existing["current_json"])
except json.JSONDecodeError as exc:
logger.warning(
"failed to parse JSON in entities.current_json: %s",
exc,
)
# 合并 current (新值覆盖旧值)
merged_current = {**old_current, **entity.current}
if update_metadata:
# 完整更新(包括元数据)
cursor.execute(
"""
UPDATE entities SET
canonical_name = ?,
tier = ?,
desc = ?,
current_json = ?,
last_appearance = ?,
is_protagonist = ?,
is_archived = ?,
updated_at = CURRENT_TIMESTAMP
WHERE id = ?
""",
(
entity.canonical_name,
entity.tier,
entity.desc,
json.dumps(merged_current, ensure_ascii=False),
entity.last_appearance,
1 if entity.is_protagonist else 0,
1 if entity.is_archived else 0,
entity.id,
),
)
else:
# 只更新 current 和 last_appearance
cursor.execute(
"""
UPDATE entities SET
current_json = ?,
last_appearance = ?,
updated_at = CURRENT_TIMESTAMP
WHERE id = ?
""",
(
json.dumps(merged_current, ensure_ascii=False),
entity.last_appearance,
entity.id,
),
)
conn.commit()
return False
else:
# 新实体: 插入
cursor.execute(
"""
INSERT INTO entities
(id, type, canonical_name, tier, desc, current_json,
first_appearance, last_appearance, is_protagonist, is_archived)
VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?)
""",
(
entity.id,
entity.type,
entity.canonical_name,
entity.tier,
entity.desc,
json.dumps(entity.current, ensure_ascii=False),
entity.first_appearance,
entity.last_appearance,
1 if entity.is_protagonist else 0,
1 if entity.is_archived else 0,
),
)
conn.commit()
return True
def get_entity(self, entity_id: str) -> Optional[Dict]:
"""获取单个实体"""
with self._get_conn() as conn:
cursor = conn.cursor()
cursor.execute("SELECT * FROM entities WHERE id = ?", (entity_id,))
row = cursor.fetchone()
if row:
return self._row_to_dict(row, parse_json=["current_json"])
return None
def get_entities_by_type(
self, entity_type: str, include_archived: bool = False
) -> List[Dict]:
"""按类型获取实体"""
with self._get_conn() as conn:
cursor = conn.cursor()
if include_archived:
cursor.execute(
"""
SELECT * FROM entities WHERE type = ?
ORDER BY last_appearance DESC
""",
(entity_type,),
)
else:
cursor.execute(
"""
SELECT * FROM entities WHERE type = ? AND is_archived = 0
ORDER BY last_appearance DESC
""",
(entity_type,),
)
return [
self._row_to_dict(row, parse_json=["current_json"])
for row in cursor.fetchall()
]
def get_entities_by_tier(self, tier: str) -> List[Dict]:
"""按重要度获取实体 (核心/重要/次要/装饰)"""
with self._get_conn() as conn:
cursor = conn.cursor()
cursor.execute(
"""
SELECT * FROM entities WHERE tier = ? AND is_archived = 0
ORDER BY last_appearance DESC
""",
(tier,),
)
return [
self._row_to_dict(row, parse_json=["current_json"])
for row in cursor.fetchall()
]
def get_core_entities(self) -> List[Dict]:
"""获取所有核心实体 (用于 Context Agent 全量加载)"""
with self._get_conn() as conn:
cursor = conn.cursor()
cursor.execute("""
SELECT * FROM entities
WHERE (tier IN ('核心', '重要') OR is_protagonist = 1) AND is_archived = 0
ORDER BY is_protagonist DESC, tier, last_appearance DESC
""")
return [
self._row_to_dict(row, parse_json=["current_json"])
for row in cursor.fetchall()
]
def get_protagonist(self) -> Optional[Dict]:
"""获取主角实体"""
with self._get_conn() as conn:
cursor = conn.cursor()
cursor.execute("SELECT * FROM entities WHERE is_protagonist = 1 LIMIT 1")
row = cursor.fetchone()
if row:
return self._row_to_dict(row, parse_json=["current_json"])
return None
def update_entity_current(self, entity_id: str, updates: Dict) -> bool:
"""
增量更新实体的 current 字段 (不覆盖其他字段)
例如: update_entity_current("xiaoyan", {"realm": "斗师"})
"""
with self._get_conn() as conn:
cursor = conn.cursor()
cursor.execute(
"SELECT current_json FROM entities WHERE id = ?", (entity_id,)
)
row = cursor.fetchone()
if not row:
return False
current = {}
if row["current_json"]:
try:
current = json.loads(row["current_json"])
except json.JSONDecodeError as exc:
logger.warning(
"failed to parse JSON in update_entity_current current_json: %s",
exc,
)
current.update(updates)
cursor.execute(
"""
UPDATE entities SET
current_json = ?,
updated_at = CURRENT_TIMESTAMP
WHERE id = ?
""",
(json.dumps(current, ensure_ascii=False), entity_id),
)
conn.commit()
return True
def archive_entity(self, entity_id: str) -> bool:
"""归档实体 (不删除,只是标记)"""
with self._get_conn() as conn:
cursor = conn.cursor()
cursor.execute(
"""
UPDATE entities SET is_archived = 1, updated_at = CURRENT_TIMESTAMP
WHERE id = ?
""",
(entity_id,),
)
conn.commit()
return cursor.rowcount > 0
# ==================== v5.1 别名操作 ====================
def register_alias(self, alias: str, entity_id: str, entity_type: str) -> bool:
"""
注册别名 (支持一对多)
同一别名可映射多个实体 (如 "天云宗" → 地点 + 势力)
"""
with self._get_conn() as conn:
cursor = conn.cursor()
try:
cursor.execute(
"""
INSERT OR IGNORE INTO aliases (alias, entity_id, entity_type)
VALUES (?, ?, ?)
""",
(alias, entity_id, entity_type),
)
conn.commit()
return cursor.rowcount > 0
except sqlite3.IntegrityError:
return False
def get_entities_by_alias(self, alias: str) -> List[Dict]:
"""
根据别名查找实体 (一对多)
返回所有匹配的实体 (可能有多个不同类型)
"""
with self._get_conn() as conn:
cursor = conn.cursor()
cursor.execute(
"""
SELECT e.*, a.entity_type as alias_type
FROM entities e
JOIN aliases a ON e.id = a.entity_id
WHERE a.alias = ?
""",
(alias,),
)
return [
self._row_to_dict(row, parse_json=["current_json"])
for row in cursor.fetchall()
]
def get_entity_aliases(self, entity_id: str) -> List[str]:
"""获取实体的所有别名"""
with self._get_conn() as conn:
cursor = conn.cursor()
cursor.execute(
"SELECT alias FROM aliases WHERE entity_id = ?", (entity_id,)
)
return [row["alias"] for row in cursor.fetchall()]
def remove_alias(self, alias: str, entity_id: str) -> bool:
"""移除别名"""
with self._get_conn() as conn:
cursor = conn.cursor()
cursor.execute(
"DELETE FROM aliases WHERE alias = ? AND entity_id = ?",
(alias, entity_id),
)
conn.commit()
return cursor.rowcount > 0
# ==================== v5.1 状态变化操作 ====================
def record_state_change(self, change: StateChangeMeta) -> int:
"""
记录状态变化
返回记录 ID
"""
with self._get_conn() as conn:
cursor = conn.cursor()
cursor.execute(
"""
INSERT INTO state_changes
(entity_id, field, old_value, new_value, reason, chapter)
VALUES (?, ?, ?, ?, ?, ?)
""",
(
change.entity_id,
change.field,
change.old_value,
change.new_value,
change.reason,
change.chapter,
),
)
conn.commit()
return cursor.lastrowid
def get_entity_state_changes(self, entity_id: str, limit: int = 20) -> List[Dict]:
"""获取实体的状态变化历史"""
with self._get_conn() as conn:
cursor = conn.cursor()
cursor.execute(
"""
SELECT * FROM state_changes
WHERE entity_id = ?
ORDER BY chapter DESC, id DESC
LIMIT ?
""",
(entity_id, limit),
)
return [dict(row) for row in cursor.fetchall()]
def get_recent_state_changes(self, limit: int = 50) -> List[Dict]:
"""获取最近的状态变化"""
with self._get_conn() as conn:
cursor = conn.cursor()
cursor.execute(
"""
SELECT * FROM state_changes
ORDER BY chapter DESC, id DESC
LIMIT ?
""",
(limit,),
)
return [dict(row) for row in cursor.fetchall()]
def get_chapter_state_changes(self, chapter: int) -> List[Dict]:
"""获取某章的所有状态变化"""
with self._get_conn() as conn:
cursor = conn.cursor()
cursor.execute(
"""
SELECT * FROM state_changes
WHERE chapter = ?
ORDER BY id
""",
(chapter,),
)
return [dict(row) for row in cursor.fetchall()]
# ==================== v5.1 关系操作 ====================
def upsert_relationship(self, rel: RelationshipMeta) -> bool:
"""
插入或更新关系
相同 (from, to, type) 会更新 description 和 chapter
返回是否为新关系
"""
with self._get_conn() as conn:
cursor = conn.cursor()
# 检查是否存在
cursor.execute(
"""
SELECT id FROM relationships
WHERE from_entity = ? AND to_entity = ? AND type = ?
""",
(rel.from_entity, rel.to_entity, rel.type),
)
existing = cursor.fetchone()
if existing:
cursor.execute(
"""
UPDATE relationships SET
description = ?,
chapter = ?
WHERE id = ?
""",
(rel.description, rel.chapter, existing["id"]),
)
conn.commit()
return False
else:
cursor.execute(
"""
INSERT INTO relationships
(from_entity, to_entity, type, description, chapter)
VALUES (?, ?, ?, ?, ?)
""",
(
rel.from_entity,
rel.to_entity,
rel.type,
rel.description,
rel.chapter,
),
)
conn.commit()
return True
def get_entity_relationships(
self, entity_id: str, direction: str = "both"
) -> List[Dict]:
"""
获取实体的关系
direction: "from" | "to" | "both"
"""
with self._get_conn() as conn:
cursor = conn.cursor()
if direction == "from":
cursor.execute(
"""
SELECT * FROM relationships WHERE from_entity = ?
ORDER BY chapter DESC
""",
(entity_id,),
)
elif direction == "to":
cursor.execute(
"""
SELECT * FROM relationships WHERE to_entity = ?
ORDER BY chapter DESC
""",
(entity_id,),
)
else: # both
cursor.execute(
"""
SELECT * FROM relationships
WHERE from_entity = ? OR to_entity = ?
ORDER BY chapter DESC
""",
(entity_id, entity_id),
)
return [dict(row) for row in cursor.fetchall()]
def get_relationship_between(self, entity1: str, entity2: str) -> List[Dict]:
"""获取两个实体之间的所有关系"""
with self._get_conn() as conn:
cursor = conn.cursor()
cursor.execute(
"""
SELECT * FROM relationships
WHERE (from_entity = ? AND to_entity = ?)
OR (from_entity = ? AND to_entity = ?)
ORDER BY chapter DESC
""",
(entity1, entity2, entity2, entity1),
)
return [dict(row) for row in cursor.fetchall()]
def get_recent_relationships(self, limit: int = 30) -> List[Dict]:
"""获取最近建立的关系"""
with self._get_conn() as conn:
cursor = conn.cursor()
cursor.execute(
"""
SELECT * FROM relationships
ORDER BY chapter DESC, id DESC
LIMIT ?
""",
(limit,),
)
return [dict(row) for row in cursor.fetchall()]
# ==================== v5.5 关系事件与图谱 ====================
def _infer_relationship_polarity(self, rel_type: str) -> int:
"""基于关系类型推断极性:-1 敌对,0 中立,1 友好。"""
t = str(rel_type or "")
positive_keywords = ("盟友", "友好", "师徒", "同伴", "亲", "爱", "合作")
negative_keywords = ("敌", "仇", "恨", "对立", "冲突", "背叛", "追杀")
if any(k in t for k in negative_keywords):
return -1
if any(k in t for k in positive_keywords):
return 1
return 0
def record_relationship_event(self, event: RelationshipEventMeta) -> int:
"""记录关系事件,返回事件 ID。"""
from_entity = str(getattr(event, "from_entity", "") or "").strip()
to_entity = str(getattr(event, "to_entity", "") or "").strip()
rel_type = str(getattr(event, "type", "") or "").strip()
if not from_entity or not to_entity or not rel_type:
return 0
action = str(getattr(event, "action", "update") or "update").strip().lower()
if action not in {"create", "update", "decay", "remove"}:
action = "update"
try:
chapter = int(getattr(event, "chapter", 0) or 0)
except (TypeError, ValueError):
return 0
if chapter <= 0:
return 0
try:
scene_index = int(getattr(event, "scene_index", 0) or 0)
except (TypeError, ValueError):
scene_index = 0
raw_polarity = getattr(event, "polarity", None)
if raw_polarity is None:
polarity = self._infer_relationship_polarity(rel_type)
else:
try:
polarity = int(raw_polarity)
except (TypeError, ValueError):
polarity = 0
if polarity > 1:
polarity = 1
elif polarity < -1:
polarity = -1
try:
strength = float(getattr(event, "strength", 0.5) or 0.5)
except (TypeError, ValueError):
strength = 0.5
strength = max(0.0, min(1.0, strength))
description = str(getattr(event, "description", "") or "").strip()
evidence = str(getattr(event, "evidence", "") or "").strip()
try:
confidence = float(getattr(event, "confidence", 1.0) or 1.0)
except (TypeError, ValueError):
confidence = 1.0
confidence = max(0.0, min(1.0, confidence))
with self._get_conn() as conn:
cursor = conn.cursor()
cursor.execute(
"""
INSERT INTO relationship_events
(from_entity, to_entity, type, action, polarity, strength, description, chapter, scene_index, evidence, confidence)
VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)
""",
(
from_entity,
to_entity,
rel_type,
action,
polarity,
strength,
description,
chapter,
scene_index,
evidence,
confidence,
),
)
conn.commit()
return int(cursor.lastrowid or 0)
def get_relationship_events(
self,
entity_id: str,
direction: str = "both",
from_chapter: Optional[int] = None,
to_chapter: Optional[int] = None,
limit: int = 100,
) -> List[Dict[str, Any]]:
"""按实体查询关系事件。"""
direction = str(direction or "both").lower()
clauses: List[str] = []
params: List[Any] = []
if direction == "from":
clauses.append("from_entity = ?")
params.append(entity_id)
elif direction == "to":
clauses.append("to_entity = ?")
params.append(entity_id)
else:
clauses.append("(from_entity = ? OR to_entity = ?)")
params.extend([entity_id, entity_id])
if from_chapter is not None:
clauses.append("chapter >= ?")
params.append(int(from_chapter))
if to_chapter is not None:
clauses.append("chapter <= ?")
params.append(int(to_chapter))
where_sql = " AND ".join(clauses) if clauses else "1=1"
with self._get_conn() as conn:
cursor = conn.cursor()
cursor.execute(
f"""
SELECT * FROM relationship_events
WHERE {where_sql}
ORDER BY chapter DESC, id DESC
LIMIT ?
""",
(*params, int(limit)),
)
return [dict(row) for row in cursor.fetchall()]
def get_relationship_timeline(
self,
entity1: str,
entity2: str,
from_chapter: Optional[int] = None,
to_chapter: Optional[int] = None,
limit: int = 100,
) -> List[Dict[str, Any]]:
"""查询两个实体之间的关系时间线。"""
clauses = [
"((from_entity = ? AND to_entity = ?) OR (from_entity = ? AND to_entity = ?))"
]
params: List[Any] = [entity1, entity2, entity2, entity1]
if from_chapter is not None:
clauses.append("chapter >= ?")
params.append(int(from_chapter))
if to_chapter is not None:
clauses.append("chapter <= ?")
params.append(int(to_chapter))
where_sql = " AND ".join(clauses)
with self._get_conn() as conn:
cursor = conn.cursor()
cursor.execute(
f"""
SELECT * FROM relationship_events
WHERE {where_sql}
ORDER BY chapter ASC, id ASC
LIMIT ?
""",
(*params, int(limit)),
)
return [dict(row) for row in cursor.fetchall()]
def _load_effective_relationship_edges(
self,
chapter: Optional[int] = None,
relation_types: Optional[List[str]] = None,
) -> List[Dict[str, Any]]:
"""加载指定章节截面的有效关系边。"""
relation_types = [str(t) for t in (relation_types or []) if str(t).strip()]
with self._get_conn() as conn:
cursor = conn.cursor()
if chapter is None:
clauses = []
params: List[Any] = []
if relation_types:
placeholders = ",".join("?" for _ in relation_types)
clauses.append(f"type IN ({placeholders})")
params.extend(relation_types)
where_sql = f"WHERE {' AND '.join(clauses)}" if clauses else ""
cursor.execute(
f"""
SELECT from_entity, to_entity, type, description, chapter
FROM relationships
{where_sql}
ORDER BY chapter DESC, id DESC
""",
tuple(params),
)
rows = cursor.fetchall()
return [
{
"from": str(r["from_entity"]),
"to": str(r["to_entity"]),
"type": str(r["type"]),
"description": str(r["description"] or ""),
"chapter": int(r["chapter"] or 0),
"action": "snapshot",
"polarity": self._infer_relationship_polarity(str(r["type"])),
"strength": 0.5,
"evidence": "",
"confidence": 1.0,
}
for r in rows
]
clauses = ["chapter <= ?"]
params = [int(chapter)]
if relation_types:
placeholders = ",".join("?" for _ in relation_types)
clauses.append(f"type IN ({placeholders})")
params.extend(relation_types)
cursor.execute(
f"""
SELECT *
FROM relationship_events
WHERE {' AND '.join(clauses)}
ORDER BY chapter DESC, id DESC
""",
tuple(params),
)
event_rows = cursor.fetchall()
# 兼容旧数据:若事件流不完整,回退 relationships 快照补边
snapshot_clauses = ["chapter <= ?"]
snapshot_params: List[Any] = [int(chapter)]
if relation_types:
placeholders = ",".join("?" for _ in relation_types)
snapshot_clauses.append(f"type IN ({placeholders})")
snapshot_params.extend(relation_types)
cursor.execute(
f"""
SELECT from_entity, to_entity, type, description, chapter
FROM relationships
WHERE {' AND '.join(snapshot_clauses)}
ORDER BY chapter DESC, id DESC
""",
tuple(snapshot_params),
)
snapshot_rows = cursor.fetchall()
# 章节截面:相同关系只保留“最近一次事件”,remove 视为已失效。
effective: List[Dict[str, Any]] = []
seen: set[tuple[str, str, str]] = set()
for row in event_rows:
key = (
str(row["from_entity"]),
str(row["to_entity"]),
str(row["type"]),
)
if key in seen:
continue
seen.add(key)
action = str(row["action"] or "update")
if action == "remove":
continue
effective.append(
{
"from": key[0],
"to": key[1],
"type": key[2],
"description": str(row["description"] or ""),
"chapter": int(row["chapter"] or 0),
"action": action,
"polarity": int(row["polarity"] or 0),
"strength": float(row["strength"] or 0.5),
"evidence": str(row["evidence"] or ""),
"confidence": float(row["confidence"] or 1.0),
}
)
# 事件流缺失时,从关系快照补齐(若 key 已出现则以事件为准)
for row in snapshot_rows:
key = (
str(row["from_entity"]),
str(row["to_entity"]),
str(row["type"]),
)
if key in seen:
continue
effective.append(
{
"from": key[0],
"to": key[1],
"type": key[2],
"description": str(row["description"] or ""),
"chapter": int(row["chapter"] or 0),
"action": "snapshot",
"polarity": self._infer_relationship_polarity(key[2]),
"strength": 0.5,
"evidence": "",
"confidence": 1.0,
}
)
return effective
def build_relationship_subgraph(
self,
center_entity: str,
depth: int = 2,
chapter: Optional[int] = None,
top_edges: int = 50,
relation_types: Optional[List[str]] = None,
) -> Dict[str, Any]:
"""按中心实体构建关系子图。"""
center_entity = str(center_entity or "").strip()
depth = max(1, int(depth or 1))
top_edges = max(1, int(top_edges or 1))
edges_all = self._load_effective_relationship_edges(
chapter=chapter,
relation_types=relation_types,
)
edges_all.sort(key=lambda x: int(x.get("chapter", 0)), reverse=True)
selected_edges: List[Dict[str, Any]] = []
selected_keys: set[tuple[str, str, str]] = set()
visited_nodes: set[str] = {center_entity} if center_entity else set()
frontier: set[str] = {center_entity} if center_entity else set()
for _ in range(depth):
if not frontier:
break
next_frontier: set[str] = set()
for edge in edges_all:
from_entity = str(edge.get("from") or "")
to_entity = str(edge.get("to") or "")
if from_entity not in frontier and to_entity not in frontier:
continue
key = (from_entity, to_entity, str(edge.get("type") or ""))
if key in selected_keys:
continue
selected_keys.add(key)
selected_edges.append(edge)
if from_entity and from_entity not in visited_nodes:
visited_nodes.add(from_entity)
next_frontier.add(from_entity)
if to_entity and to_entity not in visited_nodes:
visited_nodes.add(to_entity)
next_frontier.add(to_entity)
if len(selected_edges) >= top_edges:
break
frontier = next_frontier
if len(selected_edges) >= top_edges:
break
if center_entity and center_entity not in visited_nodes:
visited_nodes.add(center_entity)
# 查询节点详情
entity_map: Dict[str, Dict[str, Any]] = {}
if visited_nodes:
with self._get_conn() as conn:
cursor = conn.cursor()
placeholders = ",".join("?" for _ in visited_nodes)
cursor.execute(
f"""
SELECT id, canonical_name, type, tier, last_appearance
FROM entities
WHERE id IN ({placeholders})
""",
tuple(visited_nodes),
)
for row in cursor.fetchall():
entity_map[str(row["id"])] = {
"id": str(row["id"]),
"name": str(row["canonical_name"] or row["id"]),
"type": str(row["type"] or "未知"),
"tier": str(row["tier"] or "装饰"),
"last_appearance": int(row["last_appearance"] or 0),
}
nodes: List[Dict[str, Any]] = []
for entity_id in sorted(
visited_nodes,
key=lambda eid: (
0 if eid == center_entity else 1,
-(entity_map.get(eid, {}).get("last_appearance", 0)),
eid,
),
):
if entity_id in entity_map:
nodes.append(entity_map[entity_id])
else:
nodes.append(
{
"id": entity_id,
"name": entity_id or "未知",
"type": "未知",
"tier": "装饰",
"last_appearance": 0,
}
)
return {
"center": center_entity,
"depth": depth,
"chapter": chapter,
"nodes": nodes,
"edges": selected_edges[:top_edges],
"generated_at": datetime.now().isoformat(timespec="seconds"),
}
def _sanitize_mermaid_node_id(self, raw_id: str) -> str:
safe = re.sub(r"[^0-9a-zA-Z_]", "_", str(raw_id or "node"))
if not safe:
safe = "node"
if safe[0].isdigit():
safe = f"n_{safe}"
return safe
def render_relationship_subgraph_mermaid(self, graph: Dict[str, Any]) -> str:
"""将关系子图渲染为 Mermaid。"""
lines = ["```mermaid", "graph LR"]
nodes = graph.get("nodes") or []
edges = graph.get("edges") or []
if not nodes:
lines.append(" EMPTY[暂无关系数据]")
lines.append("```")
return "\n".join(lines)
node_alias: Dict[str, str] = {}
for node in nodes:
entity_id = str(node.get("id") or "")
if not entity_id:
continue
alias = self._sanitize_mermaid_node_id(entity_id)
node_alias[entity_id] = alias
label = str(node.get("name") or entity_id).replace('"', "'")
lines.append(f' {alias}["{label}"]')
for edge in edges:
from_entity = str(edge.get("from") or "")
to_entity = str(edge.get("to") or "")
if from_entity not in node_alias or to_entity not in node_alias:
continue
edge_type = str(edge.get("type") or "关联")
chapter = edge.get("chapter")
chapter_suffix = f"@{chapter}" if chapter not in (None, "") else ""
label = f"{edge_type}{chapter_suffix}".replace('"', "'")
try:
polarity = int(edge.get("polarity", 0) or 0)
except (TypeError, ValueError):
polarity = 0
if polarity < 0:
connector = "-.->"
else:
connector = "-->"
lines.append(
f" {node_alias[from_entity]} {connector}|{label}| {node_alias[to_entity]}"
)
lines.append("```")
return "\n".join(lines)
# ==================== v5.3 Override Contract 操作 ====================
def update_entity_field(self, entity_id: str, field: str, value: Any) -> bool:
"""Compatibility helper to update a single entity field in current_json."""
return self.update_entity_current(entity_id, {field: value})
File diff suppressed because it is too large. Load diff
@@ -0,0 +1,227 @@
#!/usr/bin/env python3
# -*- coding: utf-8 -*-
"""
IndexObservabilityMixin extracted from IndexManager.
"""
from __future__ import annotations
import json
import logging
from datetime import datetime
from typing import Any, Dict, List, Optional
logger = logging.getLogger(__name__)
class IndexObservabilityMixin:
def _row_to_dict(self, row: sqlite3.Row, parse_json: List[str] = None) -> Dict:
"""将 Row 转换为字典"""
d = dict(row)
if parse_json:
for key in parse_json:
if key in d and d[key]:
try:
d[key] = json.loads(d[key])
except json.JSONDecodeError as exc:
logger.warning(
"failed to parse JSON field %s in _row_to_dict: %s",
key,
exc,
)
return d
# ==================== 无效事实管理 ====================
def mark_invalid_fact(
self,
source_type: str,
source_id: str,
reason: str,
marked_by: str = "user",
chapter_discovered: Optional[int] = None,
) -> int:
"""标记无效事实(pending)"""
with self._get_conn() as conn:
cursor = conn.cursor()
cursor.execute(
"""
INSERT INTO invalid_facts
(source_type, source_id, reason, status, marked_by, chapter_discovered)
VALUES (?, ?, ?, 'pending', ?, ?)
""",
(source_type, str(source_id), reason, marked_by, chapter_discovered),
)
conn.commit()
return int(cursor.lastrowid)
def resolve_invalid_fact(self, invalid_id: int, action: str) -> bool:
"""确认或撤销无效标记"""
action = action.lower()
with self._get_conn() as conn:
cursor = conn.cursor()
if action == "confirm":
cursor.execute(
"""
UPDATE invalid_facts
SET status = 'confirmed', confirmed_at = CURRENT_TIMESTAMP
WHERE id = ?
""",
(invalid_id,),
)
elif action == "dismiss":
cursor.execute("DELETE FROM invalid_facts WHERE id = ?", (invalid_id,))
else:
return False
conn.commit()
return cursor.rowcount > 0
def list_invalid_facts(self, status: Optional[str] = None) -> List[Dict]:
"""列出无效事实"""
with self._get_conn() as conn:
cursor = conn.cursor()
if status:
cursor.execute(
"SELECT * FROM invalid_facts WHERE status = ? ORDER BY id DESC",
(status,),
)
else:
cursor.execute("SELECT * FROM invalid_facts ORDER BY id DESC")
return [dict(r) for r in cursor.fetchall()]
def get_invalid_ids(self, source_type: str, status: str = "confirmed") -> set[str]:
"""获取无效事实 ID 集合"""
with self._get_conn() as conn:
cursor = conn.cursor()
cursor.execute(
"SELECT source_id FROM invalid_facts WHERE source_type = ? AND status = ?",
(source_type, status),
)
return {str(r[0]) for r in cursor.fetchall() if r and r[0] is not None}
# ==================== 日志记录 ====================
def log_rag_query(
self,
query: str,
query_type: str,
results_count: int,
hit_sources: Optional[str] = None,
latency_ms: Optional[int] = None,
chapter: Optional[int] = None,
) -> None:
with self._get_conn() as conn:
cursor = conn.cursor()
cursor.execute(
"""
INSERT INTO rag_query_log
(query, query_type, results_count, hit_sources, latency_ms, chapter)
VALUES (?, ?, ?, ?, ?, ?)
""",
(query, query_type, results_count, hit_sources, latency_ms, chapter),
)
conn.commit()
def log_tool_call(
self,
tool_name: str,
success: bool,
retry_count: int = 0,
error_code: Optional[str] = None,
error_message: Optional[str] = None,
chapter: Optional[int] = None,
) -> None:
with self._get_conn() as conn:
cursor = conn.cursor()
cursor.execute(
"""
INSERT INTO tool_call_stats
(tool_name, success, retry_count, error_code, error_message, chapter)
VALUES (?, ?, ?, ?, ?, ?)
""",
(tool_name, int(bool(success)), retry_count, error_code, error_message, chapter),
)
conn.commit()
def get_stats(self) -> Dict[str, int]:
"""获取索引统计"""
with self._get_conn() as conn:
cursor = conn.cursor()
cursor.execute("SELECT COUNT(*) FROM chapters")
chapters = cursor.fetchone()[0]
cursor.execute("SELECT COUNT(*) FROM scenes")
scenes = cursor.fetchone()[0]
cursor.execute("SELECT COUNT(DISTINCT entity_id) FROM appearances")
appearances = cursor.fetchone()[0]
cursor.execute("SELECT MAX(chapter) FROM chapters")
max_chapter = cursor.fetchone()[0] or 0
# v5.1 引入统计
cursor.execute("SELECT COUNT(*) FROM entities")
entities = cursor.fetchone()[0]
cursor.execute("SELECT COUNT(*) FROM entities WHERE is_archived = 0")
active_entities = cursor.fetchone()[0]
cursor.execute("SELECT COUNT(*) FROM aliases")
aliases = cursor.fetchone()[0]
cursor.execute("SELECT COUNT(*) FROM state_changes")
state_changes = cursor.fetchone()[0]
cursor.execute("SELECT COUNT(*) FROM relationships")
relationships = cursor.fetchone()[0]
cursor.execute("SELECT COUNT(*) FROM relationship_events")
relationship_events = cursor.fetchone()[0]
# v5.3 引入统计
cursor.execute("SELECT COUNT(*) FROM override_contracts")
override_contracts = cursor.fetchone()[0]
cursor.execute(
"SELECT COUNT(*) FROM override_contracts WHERE status = 'pending'"
)
pending_overrides = cursor.fetchone()[0]
cursor.execute("SELECT COUNT(*) FROM chase_debt WHERE status = 'active'")
active_debts = cursor.fetchone()[0]
cursor.execute(
"SELECT COALESCE(SUM(current_amount), 0) FROM chase_debt WHERE status IN ('active', 'overdue')"
)
total_debt = cursor.fetchone()[0]
cursor.execute("SELECT COUNT(*) FROM chapter_reading_power")
reading_power_records = cursor.fetchone()[0]
cursor.execute("SELECT COUNT(*) FROM review_metrics")
review_metrics = cursor.fetchone()[0]
return {
"chapters": chapters,
"scenes": scenes,
"appearances": appearances,
"max_chapter": max_chapter,
# v5.1 引入
"entities": entities,
"active_entities": active_entities,
"aliases": aliases,
"state_changes": state_changes,
"relationships": relationships,
"relationship_events": relationship_events,
# v5.3 引入
"override_contracts": override_contracts,
"pending_overrides": pending_overrides,
"active_debts": active_debts,
"total_debt": total_debt,
"reading_power_records": reading_power_records,
"review_metrics": review_metrics,
}
@@ -0,0 +1,382 @@
#!/usr/bin/env python3
# -*- coding: utf-8 -*-
"""
IndexReadingMixin extracted from IndexManager.
"""
from __future__ import annotations
import json
import sys
from datetime import datetime
from typing import Any, Dict, List, Optional
class IndexReadingMixin:
def save_chapter_reading_power(self, meta: ChapterReadingPowerMeta):
"""保存章节追读力元数据"""
with self._get_conn() as conn:
cursor = conn.cursor()
cursor.execute(
"""
INSERT OR REPLACE INTO chapter_reading_power
(chapter, hook_type, hook_strength, coolpoint_patterns,
micropayoffs, hard_violations, soft_suggestions,
is_transition, override_count, debt_balance)
VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?)
""",
(
meta.chapter,
meta.hook_type,
meta.hook_strength,
json.dumps(meta.coolpoint_patterns, ensure_ascii=False),
json.dumps(meta.micropayoffs, ensure_ascii=False),
json.dumps(meta.hard_violations, ensure_ascii=False),
json.dumps(meta.soft_suggestions, ensure_ascii=False),
1 if meta.is_transition else 0,
meta.override_count,
meta.debt_balance,
),
)
conn.commit()
def get_chapter_reading_power(self, chapter: int) -> Optional[Dict]:
"""获取章节追读力元数据"""
with self._get_conn() as conn:
cursor = conn.cursor()
cursor.execute(
"SELECT * FROM chapter_reading_power WHERE chapter = ?", (chapter,)
)
row = cursor.fetchone()
if row:
return self._row_to_dict(
row,
parse_json=[
"coolpoint_patterns",
"micropayoffs",
"hard_violations",
"soft_suggestions",
],
)
return None
def get_recent_reading_power(self, limit: int = 10) -> List[Dict]:
"""获取最近章节的追读力元数据"""
with self._get_conn() as conn:
cursor = conn.cursor()
cursor.execute(
"""
SELECT * FROM chapter_reading_power
ORDER BY chapter DESC
LIMIT ?
""",
(limit,),
)
return [
self._row_to_dict(
row,
parse_json=[
"coolpoint_patterns",
"micropayoffs",
"hard_violations",
"soft_suggestions",
],
)
for row in cursor.fetchall()
]
def get_pattern_usage_stats(self, last_n_chapters: int = 20) -> Dict[str, int]:
"""获取最近N章的爽点模式使用统计"""
with self._get_conn() as conn:
cursor = conn.cursor()
cursor.execute(
"""
SELECT coolpoint_patterns FROM chapter_reading_power
ORDER BY chapter DESC
LIMIT ?
""",
(last_n_chapters,),
)
stats = {}
for row in cursor.fetchall():
if row["coolpoint_patterns"]:
try:
patterns = json.loads(row["coolpoint_patterns"])
for p in patterns:
stats[p] = stats.get(p, 0) + 1
except json.JSONDecodeError as exc:
print(
f"[index_manager] failed to parse JSON in chapter_reading_power.coolpoint_patterns: {exc}",
file=sys.stderr,
)
return stats
def get_hook_type_stats(self, last_n_chapters: int = 20) -> Dict[str, int]:
"""获取最近N章的钩子类型使用统计"""
with self._get_conn() as conn:
cursor = conn.cursor()
cursor.execute(
"""
SELECT hook_type FROM chapter_reading_power
WHERE hook_type IS NOT NULL AND hook_type != ''
ORDER BY chapter DESC
LIMIT ?
""",
(last_n_chapters,),
)
stats = {}
for row in cursor.fetchall():
hook = row["hook_type"]
stats[hook] = stats.get(hook, 0) + 1
return stats
# ==================== v5.4 审查指标 ====================
def save_review_metrics(self, metrics: ReviewMetrics) -> None:
"""保存审查指标记录"""
with self._get_conn() as conn:
cursor = conn.cursor()
cursor.execute(
"""
INSERT INTO review_metrics
(start_chapter, end_chapter, overall_score, dimension_scores,
severity_counts, critical_issues, report_file, notes, created_at, updated_at)
VALUES (?, ?, ?, ?, ?, ?, ?, ?, CURRENT_TIMESTAMP, CURRENT_TIMESTAMP)
ON CONFLICT(start_chapter, end_chapter)
DO UPDATE SET
overall_score = excluded.overall_score,
dimension_scores = excluded.dimension_scores,
severity_counts = excluded.severity_counts,
critical_issues = excluded.critical_issues,
report_file = excluded.report_file,
notes = excluded.notes,
updated_at = CURRENT_TIMESTAMP
""",
(
metrics.start_chapter,
metrics.end_chapter,
metrics.overall_score,
json.dumps(metrics.dimension_scores, ensure_ascii=False),
json.dumps(metrics.severity_counts, ensure_ascii=False),
json.dumps(metrics.critical_issues, ensure_ascii=False),
metrics.report_file,
metrics.notes,
),
)
conn.commit()
def get_recent_review_metrics(self, limit: int = 5) -> List[Dict]:
"""获取最近审查记录"""
with self._get_conn() as conn:
cursor = conn.cursor()
cursor.execute(
"""
SELECT * FROM review_metrics
ORDER BY end_chapter DESC, start_chapter DESC
LIMIT ?
""",
(limit,),
)
return [
self._row_to_dict(
row,
parse_json=["dimension_scores", "severity_counts", "critical_issues"],
)
for row in cursor.fetchall()
]
def get_review_trend_stats(self, last_n: int = 5) -> Dict[str, Any]:
"""获取审查趋势统计"""
records = self.get_recent_review_metrics(last_n)
if not records:
return {
"count": 0,
"overall_avg": 0.0,
"dimension_avg": {},
"severity_totals": {},
"recent_ranges": [],
}
overall_scores: List[float] = []
dimension_totals: Dict[str, float] = {}
dimension_counts: Dict[str, int] = {}
severity_totals: Dict[str, int] = {}
for record in records:
score = record.get("overall_score")
if score is not None:
try:
overall_scores.append(float(score))
except (TypeError, ValueError):
pass
dimensions = record.get("dimension_scores") or {}
if isinstance(dimensions, dict):
for key, value in dimensions.items():
try:
val = float(value)
except (TypeError, ValueError):
continue
dimension_totals[key] = dimension_totals.get(key, 0.0) + val
dimension_counts[key] = dimension_counts.get(key, 0) + 1
severities = record.get("severity_counts") or {}
if isinstance(severities, dict):
for key, value in severities.items():
try:
count = int(value)
except (TypeError, ValueError):
continue
severity_totals[key] = severity_totals.get(key, 0) + count
overall_avg = round(sum(overall_scores) / len(overall_scores), 2) if overall_scores else 0.0
dimension_avg = {
key: round(dimension_totals[key] / dimension_counts[key], 2)
for key in dimension_totals
if dimension_counts.get(key, 0) > 0
}
recent_ranges = [
{
"start_chapter": record.get("start_chapter"),
"end_chapter": record.get("end_chapter"),
"overall_score": record.get("overall_score", 0),
}
for record in records
]
return {
"count": len(records),
"overall_avg": overall_avg,
"dimension_avg": dimension_avg,
"severity_totals": severity_totals,
"recent_ranges": recent_ranges,
}
# ==================== 写作清单评分(Phase F) ====================
def save_writing_checklist_score(self, meta: WritingChecklistScoreMeta) -> None:
"""保存章节写作清单评分。"""
with self._get_conn() as conn:
cursor = conn.cursor()
cursor.execute(
"""
INSERT INTO writing_checklist_scores (
chapter, template, total_items, required_items,
completed_items, completed_required,
total_weight, completed_weight, completion_rate, score,
score_breakdown, pending_items, source, notes
) VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)
ON CONFLICT(chapter) DO UPDATE SET
template=excluded.template,
total_items=excluded.total_items,
required_items=excluded.required_items,
completed_items=excluded.completed_items,
completed_required=excluded.completed_required,
total_weight=excluded.total_weight,
completed_weight=excluded.completed_weight,
completion_rate=excluded.completion_rate,
score=excluded.score,
score_breakdown=excluded.score_breakdown,
pending_items=excluded.pending_items,
source=excluded.source,
notes=excluded.notes,
updated_at=CURRENT_TIMESTAMP
""",
(
meta.chapter,
meta.template,
meta.total_items,
meta.required_items,
meta.completed_items,
meta.completed_required,
meta.total_weight,
meta.completed_weight,
meta.completion_rate,
meta.score,
json.dumps(meta.score_breakdown, ensure_ascii=False),
json.dumps(meta.pending_items, ensure_ascii=False),
meta.source,
meta.notes,
),
)
conn.commit()
def get_writing_checklist_score(self, chapter: int) -> Optional[Dict[str, Any]]:
"""获取指定章节的写作清单评分。"""
with self._get_conn() as conn:
cursor = conn.cursor()
cursor.execute(
"SELECT * FROM writing_checklist_scores WHERE chapter = ?",
(chapter,),
)
row = cursor.fetchone()
if not row:
return None
return self._row_to_dict(row, parse_json=["score_breakdown", "pending_items"])
def get_recent_writing_checklist_scores(self, limit: int = 10) -> List[Dict[str, Any]]:
"""获取最近章节写作清单评分。"""
with self._get_conn() as conn:
cursor = conn.cursor()
cursor.execute(
"""
SELECT * FROM writing_checklist_scores
ORDER BY chapter DESC
LIMIT ?
""",
(limit,),
)
return [
self._row_to_dict(row, parse_json=["score_breakdown", "pending_items"])
for row in cursor.fetchall()
]
def get_writing_checklist_score_trend(self, last_n: int = 10) -> Dict[str, Any]:
"""获取写作清单评分趋势统计。"""
records = self.get_recent_writing_checklist_scores(limit=max(1, int(last_n)))
if not records:
return {
"count": 0,
"score_avg": 0.0,
"completion_avg": 0.0,
"required_completion_avg": 0.0,
"recent": [],
}
scores: List[float] = []
completion_rates: List[float] = []
required_rates: List[float] = []
for row in records:
try:
scores.append(float(row.get("score", 0.0)))
except (TypeError, ValueError):
pass
try:
completion_rates.append(float(row.get("completion_rate", 0.0)))
except (TypeError, ValueError):
pass
required_items = int(row.get("required_items") or 0)
completed_required = int(row.get("completed_required") or 0)
if required_items > 0:
required_rates.append(completed_required / required_items)
else:
required_rates.append(1.0)
return {
"count": len(records),
"score_avg": round(sum(scores) / len(scores), 2) if scores else 0.0,
"completion_avg": round(sum(completion_rates) / len(completion_rates), 4) if completion_rates else 0.0,
"required_completion_avg": round(sum(required_rates) / len(required_rates), 4) if required_rates else 0.0,
"recent": [
{
"chapter": row.get("chapter"),
"score": row.get("score"),
"completion_rate": row.get("completion_rate"),
}
for row in records
],
}
@@ -0,0 +1,379 @@
#!/usr/bin/env python3
# -*- coding: utf-8 -*-
"""
migrate_state_to_sqlite.py - 数据迁移脚本 (v5.4)
将 state.json 中的大数据迁移到 SQLite (index.db):
- entities_v3 → entities 表
- alias_index → aliases 表
- state_changes → state_changes 表
- structured_relationships → relationships 表
迁移后 state.json 只保留精简数据 (< 5KB):
- progress
- protagonist_state
- strand_tracker
- disambiguation_warnings/pending
- project_info
- world_settings (骨架)
- plot_threads
- relationships (简化版)
- review_checkpoints
用法:
python -m data_modules.migrate_state_to_sqlite --project-root "D:/wk/斗破苍穹"
python -m data_modules.migrate_state_to_sqlite --project-root "." --dry-run
python -m data_modules.migrate_state_to_sqlite --project-root "." --backup
"""
import json
import shutil
from pathlib import Path
from datetime import datetime
from typing import Dict, Any, List
from .config import get_config, DataModulesConfig
from .sql_state_manager import SQLStateManager, EntityData
def migrate_state_to_sqlite(
config: DataModulesConfig,
dry_run: bool = False,
backup: bool = True,
verbose: bool = True
) -> Dict[str, int]:
"""
执行迁移
参数:
- config: 配置对象
- dry_run: 只分析不实际写入
- backup: 迁移前备份 state.json
- verbose: 打印详细日志
返回: 迁移统计
"""
stats = {
"entities": 0,
"aliases": 0,
"state_changes": 0,
"relationships": 0,
"skipped": 0,
"errors": 0
}
# 读取 state.json
state_file = config.state_file
if not state_file.exists():
if verbose:
print(f"❌ state.json 不存在: {state_file}")
return stats
with open(state_file, 'r', encoding='utf-8') as f:
state = json.load(f)
if verbose:
file_size = state_file.stat().st_size / 1024
print(f"📄 读取 state.json ({file_size:.1f} KB)")
# 备份
if backup and not dry_run:
backup_file = state_file.with_suffix(f".json.backup-{datetime.now().strftime('%Y%m%d_%H%M%S')}")
shutil.copy(state_file, backup_file)
if verbose:
print(f"💾 已备份到: {backup_file}")
# 初始化 SQLStateManager
sql_manager = SQLStateManager(config)
# 1. 迁移 entities_v3
entities_v3 = state.get("entities_v3", {})
if verbose:
print(f"\n🔄 迁移 entities_v3...")
for entity_type, entities in entities_v3.items():
if not isinstance(entities, dict):
continue
for entity_id, entity_data in entities.items():
if not isinstance(entity_data, dict):
stats["skipped"] += 1
continue
try:
entity = EntityData(
id=entity_id,
type=entity_type,
name=entity_data.get("canonical_name", entity_data.get("name", entity_id)),
tier=entity_data.get("tier", "装饰"),
desc=entity_data.get("desc", ""),
current=entity_data.get("current", {}),
aliases=[], # 别名单独处理
first_appearance=entity_data.get("first_appearance", 0),
last_appearance=entity_data.get("last_appearance", 0),
is_protagonist=entity_data.get("is_protagonist", False)
)
if not dry_run:
sql_manager.upsert_entity(entity)
stats["entities"] += 1
if verbose and stats["entities"] % 50 == 0:
print(f" 已迁移 {stats['entities']} 个实体...")
except Exception as e:
stats["errors"] += 1
if verbose:
print(f" ⚠️ 实体迁移失败 {entity_id}: {e}")
if verbose:
print(f" ✅ 实体: {stats['entities']} 个")
# 2. 迁移 alias_index
alias_index = state.get("alias_index", {})
if verbose:
print(f"\n🔄 迁移 alias_index...")
for alias, entries in alias_index.items():
if not isinstance(entries, list):
continue
for entry in entries:
if not isinstance(entry, dict):
stats["skipped"] += 1
continue
entity_id = entry.get("id")
entity_type = entry.get("type")
if not entity_id or not entity_type:
stats["skipped"] += 1
continue
try:
if not dry_run:
sql_manager.register_alias(alias, entity_id, entity_type)
stats["aliases"] += 1
except Exception as e:
stats["errors"] += 1
if verbose:
print(f" ⚠️ 别名迁移失败 {alias}: {e}")
if verbose:
print(f" ✅ 别名: {stats['aliases']} 个")
# 3. 迁移 state_changes
state_changes = state.get("state_changes", [])
if verbose:
print(f"\n🔄 迁移 state_changes...")
for change in state_changes:
if not isinstance(change, dict):
stats["skipped"] += 1
continue
try:
entity_id = change.get("entity_id", "")
if not entity_id:
stats["skipped"] += 1
continue
if not dry_run:
sql_manager.record_state_change(
entity_id=entity_id,
field=change.get("field", ""),
old_value=change.get("old", change.get("old_value", "")),
new_value=change.get("new", change.get("new_value", "")),
reason=change.get("reason", ""),
chapter=change.get("chapter", 0)
)
stats["state_changes"] += 1
except Exception as e:
stats["errors"] += 1
if verbose:
print(f" ⚠️ 状态变化迁移失败: {e}")
if verbose:
print(f" ✅ 状态变化: {stats['state_changes']} 条")
# 4. 迁移 structured_relationships
relationships = state.get("structured_relationships", [])
if verbose:
print(f"\n🔄 迁移 structured_relationships...")
for rel in relationships:
if not isinstance(rel, dict):
stats["skipped"] += 1
continue
try:
from_entity = rel.get("from", rel.get("from_entity", ""))
to_entity = rel.get("to", rel.get("to_entity", ""))
if not from_entity or not to_entity:
stats["skipped"] += 1
continue
if not dry_run:
sql_manager.upsert_relationship(
from_entity=from_entity,
to_entity=to_entity,
type=rel.get("type", "相识"),
description=rel.get("description", ""),
chapter=rel.get("chapter", 0)
)
stats["relationships"] += 1
except Exception as e:
stats["errors"] += 1
if verbose:
print(f" ⚠️ 关系迁移失败: {e}")
if verbose:
print(f" ✅ 关系: {stats['relationships']} 条")
# 5. 精简 state.json(移除已迁移字段)
if not dry_run:
if verbose:
print(f"\n🔄 精简 state.json...")
# 保留字段
slim_state = {
"project_info": state.get("project_info", {}),
"progress": state.get("progress", {}),
"protagonist_state": state.get("protagonist_state", {}),
"strand_tracker": state.get("strand_tracker", {}),
"world_settings": _slim_world_settings(state.get("world_settings", {})),
"plot_threads": state.get("plot_threads", {}),
"relationships": _slim_relationships(state.get("relationships", {})),
"review_checkpoints": state.get("review_checkpoints", [])[-10:], # 只保留最近10个
"disambiguation_warnings": state.get("disambiguation_warnings", [])[-20:],
"disambiguation_pending": state.get("disambiguation_pending", [])[-10:],
# v5.1 引入标记
"_migrated_to_sqlite": True,
"_migration_timestamp": datetime.now().isoformat()
}
with open(state_file, 'w', encoding='utf-8') as f:
json.dump(slim_state, f, ensure_ascii=False, indent=2)
new_size = state_file.stat().st_size / 1024
if verbose:
print(f" ✅ 精简后: {new_size:.1f} KB")
# 打印统计
if verbose:
print(f"\n" + "=" * 50)
print(f"📊 迁移统计:")
print(f" 实体: {stats['entities']}")
print(f" 别名: {stats['aliases']}")
print(f" 状态变化: {stats['state_changes']}")
print(f" 关系: {stats['relationships']}")
print(f" 跳过: {stats['skipped']}")
print(f" 错误: {stats['errors']}")
if dry_run:
print(f"\n⚠️ 这是 dry-run 模式,实际未写入任何数据")
return stats
def _slim_world_settings(world_settings: Dict) -> Dict:
"""精简 world_settings,只保留骨架"""
if not isinstance(world_settings, dict):
return {}
slim = {}
# power_system: 只保留等级名称
power_system = world_settings.get("power_system", [])
if isinstance(power_system, list):
slim["power_system"] = [
p.get("name") if isinstance(p, dict) else p
for p in power_system[:20] # 最多20个等级
]
# factions: 只保留名称和简述
factions = world_settings.get("factions", [])
if isinstance(factions, list):
slim["factions"] = [
{"name": f.get("name"), "type": f.get("type")}
if isinstance(f, dict) else f
for f in factions[:30] # 最多30个势力
]
# locations: 只保留名称
locations = world_settings.get("locations", [])
if isinstance(locations, list):
slim["locations"] = [
loc.get("name") if isinstance(loc, dict) else loc
for loc in locations[:50] # 最多50个地点
]
return slim
def _slim_relationships(relationships: Dict) -> Dict:
"""精简 relationships,只保留核心关系"""
if not isinstance(relationships, dict):
return {}
# 只保留 relationships 字典本身,不做额外精简
# 因为这个字段本身应该比较小
return relationships
def main():
import argparse
from .cli_output import print_success, print_error
from .index_manager import IndexManager
parser = argparse.ArgumentParser(description="迁移 state.json 到 SQLite (v5.4)")
parser.add_argument("--project-root", type=str, required=True, help="项目根目录")
parser.add_argument("--dry-run", action="store_true", help="只分析不实际写入")
parser.add_argument("--backup", action="store_true", default=True, help="迁移前备份")
parser.add_argument("--no-backup", action="store_true", help="不备份")
parser.add_argument("--quiet", action="store_true", help="安静模式")
args = parser.parse_args()
# 允许传入“工作区根目录”,统一解析到真正的 book project_root(必须包含 .noma/state.json)
from project_locator import resolve_project_root
resolved_root = resolve_project_root(args.project_root)
config = DataModulesConfig.from_project_root(resolved_root)
backup = not args.no_backup
logger = IndexManager(config)
tool_name = "migrate_state_to_sqlite"
try:
stats = migrate_state_to_sqlite(
config=config,
dry_run=args.dry_run,
backup=backup,
verbose=False,
)
except Exception as exc:
print_error("MIGRATE_FAILED", str(exc), suggestion="检查 state.json 与 index.db 权限")
try:
logger.log_tool_call(tool_name, False, error_code="MIGRATE_FAILED", error_message=str(exc))
except Exception:
pass
raise SystemExit(1)
if stats.get("errors", 0) > 0:
print_error("MIGRATE_ERRORS", "迁移出现错误", details=stats)
try:
logger.log_tool_call(tool_name, False, error_code="MIGRATE_ERRORS", error_message="迁移出现错误")
except Exception:
pass
raise SystemExit(1)
print_success({"project": str(config.project_root), **stats}, message="migrated")
try:
logger.log_tool_call(tool_name, True)
except Exception:
pass
if __name__ == "__main__":
main()
+318
View File
@@ -0,0 +1,318 @@
#!/usr/bin/env python3
# -*- coding: utf-8 -*-
"""
noma 统一入口(面向 skills / agents 的稳定 CLI)
设计目标:
- 只有一个入口命令,避免到处拼 `python -m data_modules.xxx ...` 导致参数位置/引号/路径炸裂。
- 自动解析正确的 book project_root(包含 `.noma/state.json` 的目录)。
- 所有写入类命令在解析到 project_root 后,统一前置 `--project-root` 传给具体模块。
典型用法(推荐,不依赖 PYTHONPATH / 不要求 cd):
python "<SCRIPTS_DIR>/noma.py" preflight
python "<SCRIPTS_DIR>/noma.py" where
python "<SCRIPTS_DIR>/noma.py" use D:\\wk\\xiaoshuo\\凡人资本论
python "<SCRIPTS_DIR>/noma.py" --project-root D:\\wk\\xiaoshuo index stats
python "<SCRIPTS_DIR>/noma.py" --project-root D:\\wk\\xiaoshuo state process-chapter --chapter 100 --data @payload.json
python "<SCRIPTS_DIR>/noma.py" --project-root D:\\wk\\xiaoshuo extract-context --chapter 100 --format json
也支持(不推荐,容易踩 PYTHONPATH/cd/参数顺序坑):
python -m data_modules.noma where
"""
from __future__ import annotations
import argparse
import importlib
import json
import subprocess
import sys
from pathlib import Path
from typing import Optional
from runtime_compat import normalize_windows_path
from project_locator import resolve_project_root, write_current_project_pointer, update_global_registry_current_project
def _scripts_dir() -> Path:
# data_modules/noma.py -> data_modules -> scripts
return Path(__file__).resolve().parent.parent
def _resolve_root(explicit_project_root: Optional[str]) -> Path:
# 允许显式传入工作区根目录或书项目根目录
raw = explicit_project_root
if raw:
return resolve_project_root(raw)
return resolve_project_root()
def _strip_project_root_args(argv: list[str]) -> list[str]:
"""
下游工具统一由本入口注入 `--project-root`,避免重复传参导致 argparse 报错/歧义。
"""
out: list[str] = []
i = 0
while i < len(argv):
tok = argv[i]
if tok == "--project-root":
i += 2
continue
if tok.startswith("--project-root="):
i += 1
continue
out.append(tok)
i += 1
return out
def _run_data_module(module: str, argv: list[str]) -> int:
"""
Import `data_modules.<module>` and call its main(), while isolating sys.argv.
"""
mod = importlib.import_module(f"data_modules.{module}")
main = getattr(mod, "main", None)
if not callable(main):
raise RuntimeError(f"data_modules.{module} 缺少可调用的 main()")
old_argv = sys.argv
try:
sys.argv = [f"data_modules.{module}"] + argv
try:
main()
return 0
except SystemExit as e:
return int(e.code or 0)
finally:
sys.argv = old_argv
def _run_script(script_name: str, argv: list[str]) -> int:
"""
Run a script under `.claude/scripts/` via a subprocess.
用途:兼容没有 main() 的脚本(例如 workflow_manager.py)。
"""
script_path = _scripts_dir() / script_name
if not script_path.is_file():
raise FileNotFoundError(f"未找到脚本: {script_path}")
proc = subprocess.run([sys.executable, str(script_path), *argv])
return int(proc.returncode or 0)
def cmd_where(args: argparse.Namespace) -> int:
root = _resolve_root(args.project_root)
print(str(root))
return 0
def _build_preflight_report(explicit_project_root: Optional[str]) -> dict:
scripts_dir = _scripts_dir().resolve()
plugin_root = scripts_dir.parent
skill_root = plugin_root / "skills" / "noma-write"
entry_script = scripts_dir / "noma.py"
extract_script = scripts_dir / "extract_chapter_context.py"
checks: list[dict[str, object]] = [
{"name": "scripts_dir", "ok": scripts_dir.is_dir(), "path": str(scripts_dir)},
{"name": "entry_script", "ok": entry_script.is_file(), "path": str(entry_script)},
{"name": "extract_context_script", "ok": extract_script.is_file(), "path": str(extract_script)},
{"name": "skill_root", "ok": skill_root.is_dir(), "path": str(skill_root)},
]
project_root = ""
project_root_error = ""
try:
resolved_root = _resolve_root(explicit_project_root)
project_root = str(resolved_root)
checks.append({"name": "project_root", "ok": True, "path": project_root})
except Exception as exc:
project_root_error = str(exc)
checks.append({"name": "project_root", "ok": False, "path": explicit_project_root or "", "error": project_root_error})
return {
"ok": all(bool(item["ok"]) for item in checks),
"project_root": project_root,
"scripts_dir": str(scripts_dir),
"skill_root": str(skill_root),
"checks": checks,
"project_root_error": project_root_error,
}
def cmd_preflight(args: argparse.Namespace) -> int:
report = _build_preflight_report(args.project_root)
if args.format == "json":
print(json.dumps(report, ensure_ascii=False, indent=2))
else:
for item in report["checks"]:
status = "OK" if item["ok"] else "ERROR"
path = item.get("path") or ""
print(f"{status} {item['name']}: {path}")
if item.get("error"):
print(f" detail: {item['error']}")
return 0 if report["ok"] else 1
def cmd_use(args: argparse.Namespace) -> int:
project_root = normalize_windows_path(args.project_root).expanduser()
try:
project_root = project_root.resolve()
except Exception:
project_root = project_root
workspace_root: Optional[Path] = None
if args.workspace_root:
workspace_root = normalize_windows_path(args.workspace_root).expanduser()
try:
workspace_root = workspace_root.resolve()
except Exception:
workspace_root = workspace_root
# 1) 写入工作区指针(若工作区内存在 `.claude/`)
pointer_file = write_current_project_pointer(project_root, workspace_root=workspace_root)
if pointer_file is not None:
print(f"workspace pointer: {pointer_file}")
else:
print("workspace pointer: (skipped)")
# 2) 写入用户级 registry(保证全局安装/空上下文可恢复)
reg_path = update_global_registry_current_project(workspace_root=workspace_root, project_root=project_root)
if reg_path is not None:
print(f"global registry: {reg_path}")
else:
print("global registry: (skipped)")
return 0
def main() -> None:
parser = argparse.ArgumentParser(description="noma unified CLI")
parser.add_argument("--project-root", help="书项目根目录或工作区根目录(可选,默认自动检测)")
sub = parser.add_subparsers(dest="tool", required=True)
p_where = sub.add_parser("where", help="打印解析出的 project_root")
p_where.set_defaults(func=cmd_where)
p_preflight = sub.add_parser("preflight", help="校验统一 CLI 运行环境与 project_root")
p_preflight.add_argument("--format", choices=["text", "json"], default="text", help="输出格式")
p_preflight.set_defaults(func=cmd_preflight)
p_use = sub.add_parser("use", help="绑定当前工作区使用的书项目(写入指针/registry)")
p_use.add_argument("project_root", help="书项目根目录(必须包含 .noma/state.json)")
p_use.add_argument("--workspace-root", help="工作区根目录(可选;默认由运行环境推断)")
p_use.set_defaults(func=cmd_use)
# Pass-through to data modules
p_index = sub.add_parser("index", help="转发到 index_manager")
p_index.add_argument("args", nargs=argparse.REMAINDER)
p_state = sub.add_parser("state", help="转发到 state_manager")
p_state.add_argument("args", nargs=argparse.REMAINDER)
p_rag = sub.add_parser("rag", help="转发到 rag_adapter")
p_rag.add_argument("args", nargs=argparse.REMAINDER)
p_style = sub.add_parser("style", help="转发到 style_sampler")
p_style.add_argument("args", nargs=argparse.REMAINDER)
p_entity = sub.add_parser("entity", help="转发到 entity_linker")
p_entity.add_argument("args", nargs=argparse.REMAINDER)
p_context = sub.add_parser("context", help="转发到 context_manager")
p_context.add_argument("args", nargs=argparse.REMAINDER)
p_migrate = sub.add_parser("migrate", help="转发到 migrate_state_to_sqlite")
p_migrate.add_argument("args", nargs=argparse.REMAINDER)
p_wiki = sub.add_parser("wiki", help="转发到 wiki_manager")
p_wiki.add_argument("args", nargs=argparse.REMAINDER)
# Pass-through to scripts
p_workflow = sub.add_parser("workflow", help="转发到 workflow_manager.py")
p_workflow.add_argument("args", nargs=argparse.REMAINDER)
p_status = sub.add_parser("status", help="转发到 status_reporter.py")
p_status.add_argument("args", nargs=argparse.REMAINDER)
p_update_state = sub.add_parser("update-state", help="转发到 update_state.py")
p_update_state.add_argument("args", nargs=argparse.REMAINDER)
p_backup = sub.add_parser("backup", help="转发到 backup_manager.py")
p_backup.add_argument("args", nargs=argparse.REMAINDER)
p_archive = sub.add_parser("archive", help="转发到 archive_manager.py")
p_archive.add_argument("args", nargs=argparse.REMAINDER)
p_init = sub.add_parser("init", help="转发到 init_project.py(初始化项目)")
p_init.add_argument("args", nargs=argparse.REMAINDER)
p_extract_context = sub.add_parser("extract-context", help="转发到 extract_chapter_context.py")
p_extract_context.add_argument("--chapter", type=int, required=True, help="目标章节号")
p_extract_context.add_argument("--format", choices=["text", "json"], default="text", help="输出格式")
# 兼容:允许 `--project-root` 出现在任意位置(减少 agents/skills 拼命令的出错率)
from .cli_args import normalize_global_project_root
argv = normalize_global_project_root(sys.argv[1:])
args = parser.parse_args(argv)
# where/use 直接执行
if hasattr(args, "func"):
code = int(args.func(args) or 0)
raise SystemExit(code)
tool = args.tool
rest = list(getattr(args, "args", []) or [])
# argparse.REMAINDER 可能以 `--` 开头占位,这里去掉
if rest[:1] == ["--"]:
rest = rest[1:]
rest = _strip_project_root_args(rest)
# init 是创建项目,不应该依赖/注入已存在 project_root
if tool == "init":
raise SystemExit(_run_script("init_project.py", rest))
# 其余工具:统一解析 project_root 后前置给下游
project_root = _resolve_root(args.project_root)
forward_args = ["--project-root", str(project_root)]
if tool == "index":
raise SystemExit(_run_data_module("index_manager", [*forward_args, *rest]))
if tool == "state":
raise SystemExit(_run_data_module("state_manager", [*forward_args, *rest]))
if tool == "rag":
raise SystemExit(_run_data_module("rag_adapter", [*forward_args, *rest]))
if tool == "style":
raise SystemExit(_run_data_module("style_sampler", [*forward_args, *rest]))
if tool == "entity":
raise SystemExit(_run_data_module("entity_linker", [*forward_args, *rest]))
if tool == "context":
raise SystemExit(_run_data_module("context_manager", [*forward_args, *rest]))
if tool == "migrate":
raise SystemExit(_run_data_module("migrate_state_to_sqlite", [*forward_args, *rest]))
if tool == "wiki":
raise SystemExit(_run_data_module("wiki_manager", [*forward_args, *rest]))
if tool == "workflow":
raise SystemExit(_run_script("workflow_manager.py", [*forward_args, *rest]))
if tool == "status":
raise SystemExit(_run_script("status_reporter.py", [*forward_args, *rest]))
if tool == "update-state":
raise SystemExit(_run_script("update_state.py", [*forward_args, *rest]))
if tool == "backup":
raise SystemExit(_run_script("backup_manager.py", [*forward_args, *rest]))
if tool == "archive":
raise SystemExit(_run_script("archive_manager.py", [*forward_args, *rest]))
if tool == "extract-context":
from runtime_compat import normalize_windows_path
chapter_path = normalize_windows_path(f"正文/第{args.chapter:04d}章.md")
return_args = [*forward_args, "--chapter", str(chapter_path)]
raise SystemExit(_run_script("extract_chapter_context.py", return_args))
raise SystemExit(2)
if __name__ == "__main__":
main()
@@ -0,0 +1,87 @@
#!/usr/bin/env python3
# -*- coding: utf-8 -*-
"""
Shared observability helpers for data modules.
"""
from __future__ import annotations
import json
import logging
from datetime import datetime
from pathlib import Path
from typing import Any, Dict, Optional
logger = logging.getLogger(__name__)
def safe_log_tool_call(
tool_logger,
*,
tool_name: str,
success: bool,
retry_count: int = 0,
error_code: Optional[str] = None,
error_message: Optional[str] = None,
chapter: Optional[int] = None,
) -> None:
try:
tool_logger.log_tool_call(
tool_name,
success,
retry_count=retry_count,
error_code=error_code,
error_message=error_message,
chapter=chapter,
)
except Exception as exc:
logger.warning(
"failed to log tool call %s: %s",
tool_name,
exc,
)
def safe_append_perf_timing(
project_root: str | Path,
*,
tool_name: str,
success: bool,
elapsed_ms: int,
chapter: Optional[int] = None,
error_code: Optional[str] = None,
error_message: Optional[str] = None,
meta: Optional[Dict[str, Any]] = None,
) -> None:
"""
Append timing trace for profiling long-running data-agent pipeline steps.
Output path:
- {project_root}/.noma/observability/data_agent_timing.jsonl
"""
try:
root = Path(project_root).resolve()
obs_dir = root / ".noma" / "observability"
obs_dir.mkdir(parents=True, exist_ok=True)
log_path = obs_dir / "data_agent_timing.jsonl"
payload: Dict[str, Any] = {
"timestamp": datetime.now().isoformat(),
"tool_name": tool_name,
"success": bool(success),
"elapsed_ms": int(max(0, elapsed_ms)),
}
if chapter is not None:
payload["chapter"] = int(chapter)
if error_code:
payload["error_code"] = error_code
if error_message:
payload["error_message"] = error_message
if meta:
payload["meta"] = meta
with open(log_path, "a", encoding="utf-8") as f:
f.write(json.dumps(payload, ensure_ascii=False) + "\n")
except Exception as exc:
logger.warning("failed to append perf timing for %s: %s", tool_name, exc)
+144
View File
@@ -0,0 +1,144 @@
#!/usr/bin/env python3
# -*- coding: utf-8 -*-
"""Query router for RAG requests."""
from __future__ import annotations
import re
from typing import Any, Dict, List
class QueryRouter:
def __init__(self):
self.intent_patterns = {
"relationship": [r"关系", r"图谱", r"时间线", r"谁和谁", r"敌对", r"盟友"],
"entity": [r"人物", r"角色", r"谁", r"身份", r"别名"],
"scene": [r"地点", r"场景", r"哪里", r"位置"],
"setting": [r"设定", r"规则", r"体系", r"世界观"],
"plot": [r"剧情", r"发生", r"事件", r"经过"],
}
self.patterns = {
"entity": list(self.intent_patterns["entity"]),
"scene": list(self.intent_patterns["scene"]),
"setting": list(self.intent_patterns["setting"]),
"plot": list(self.intent_patterns["plot"]),
}
def _extract_entities(self, query: str) -> List[str]:
# 轻量启发式提取:提取长度 2-6 的中文短语,过滤常见查询词
candidates = re.findall(r"[\u4e00-\u9fff]{2,6}", query)
stopwords = {
"关系",
"图谱",
"时间线",
"剧情",
"发生",
"事件",
"角色",
"人物",
"设定",
"世界观",
"地点",
"场景",
}
entities: List[str] = []
for c in candidates:
if c in stopwords:
continue
if c not in entities:
entities.append(c)
return entities[:4]
def _extract_time_scope(self, query: str) -> Dict[str, Any]:
m_range = re.search(r"第?\s*(\d+)\s*[-~到]\s*(\d+)\s*章", query)
if m_range:
start = int(m_range.group(1))
end = int(m_range.group(2))
if start > end:
start, end = end, start
return {"from_chapter": start, "to_chapter": end}
m_single = re.search(r"第?\s*(\d+)\s*章", query)
if m_single:
chapter = int(m_single.group(1))
return {"from_chapter": chapter, "to_chapter": chapter}
return {}
def route_intent(self, query: str) -> Dict[str, Any]:
query = str(query or "")
intent = "plot"
for intent_name, patterns in self.intent_patterns.items():
if any(re.search(pat, query) for pat in patterns):
intent = intent_name
break
time_scope = self._extract_time_scope(query)
entities = self._extract_entities(query)
needs_graph = intent == "relationship" or "关系" in query or "图谱" in query
return {
"intent": intent,
"entities": entities,
"time_scope": time_scope,
"needs_graph": needs_graph,
"raw_query": query,
}
def plan_subqueries(self, intent_payload: Dict[str, Any]) -> List[Dict[str, Any]]:
intent = str((intent_payload or {}).get("intent") or "plot")
entities = list((intent_payload or {}).get("entities") or [])
time_scope = dict((intent_payload or {}).get("time_scope") or {})
needs_graph = bool((intent_payload or {}).get("needs_graph"))
steps: List[Dict[str, Any]] = []
if intent == "relationship":
steps.append(
{
"name": "relationship_graph",
"strategy": "graph_lookup",
"entities": entities,
"time_scope": time_scope,
}
)
steps.append(
{
"name": "relationship_evidence",
"strategy": "graph_hybrid",
"entities": entities,
"time_scope": time_scope,
}
)
return steps
if needs_graph and entities:
steps.append(
{
"name": "graph_enhanced_retrieval",
"strategy": "graph_hybrid",
"entities": entities,
"time_scope": time_scope,
}
)
return steps
strategy_map = {
"entity": "hybrid",
"scene": "bm25",
"setting": "bm25",
"plot": "hybrid",
}
steps.append(
{
"name": "default_retrieval",
"strategy": strategy_map.get(intent, "hybrid"),
"entities": entities,
"time_scope": time_scope,
}
)
return steps
def route(self, query: str) -> str:
return str(self.route_intent(query).get("intent") or "plot")
def split(self, query: str) -> List[str]:
parts = re.split(r"[,,;;以及和]\s*", query)
return [p.strip() for p in parts if p.strip()]
File diff suppressed because it is too large. Load diff
+469
View File
@@ -0,0 +1,469 @@
#!/usr/bin/env python3
# -*- coding: utf-8 -*-
"""
RAG Manager - RAG 检索与管理模式
功能:
1. 检索 - 查询项目/系统/插件各层 RAG
2. 添加 - 将学习到的模式存入指定层
3. 删除 - 从指定层删除模式
4. 列表 - 列出各层的所有模式
5. 同步 - 将项目学习成果同步到系统共享层
用法:
python rag_manager.py --project-root . list --layer system
python rag_manager.py --project-root . search "打脸爽点"
python rag_manager.py --project-root . add --pattern-id xxx --layer project
python rag_manager.py --project-root . delete --pattern-id xxx --layer system
python rag_manager.py --project-root . sync --from project --to system
"""
import argparse
import asyncio
import json
import sys
from dataclasses import asdict
from datetime import datetime
from pathlib import Path
from typing import Optional
try:
from runtime_compat import enable_windows_utf8_stdio
except ImportError:
enable_windows_utf8_stdio = lambda: None
# 添加 scripts 目录到路径
sys.path.insert(0, str(Path(__file__).parent.parent))
from data_modules.cross_project_rag import CrossProjectRAG, RAGLayer, LearnedPattern
from data_modules.chapter_analyzer import ChapterAnalyzer, ChapterAnalysisResult
class RAGManager:
"""
RAG 管理器
提供:
- 检索:跨三层 RAG 检索
- 添加:存储模式到指定层
- 删除:删除指定模式
- 列表:列出各层模式
- 同步:将模式从项目层同步到系统层
"""
def __init__(self, project_root: Path):
self.project_root = project_root
self.cross_rag = CrossProjectRAG(project_root)
self.analyzer = ChapterAnalyzer(project_root)
# ==================== 检索 ====================
async def search(
self,
query: str,
top_k: int = 5,
layers: Optional[list[str]] = None
) -> list:
"""检索 RAG"""
if layers is None:
layer_list = [RAGLayer.PROJECT, RAGLayer.SYSTEM, RAGLayer.PLUGIN]
else:
layer_map = {
"project": RAGLayer.PROJECT,
"system": RAGLayer.SYSTEM,
"plugin": RAGLayer.PLUGIN
}
layer_list = [layer_map[l] for l in layers if l in layer_map]
results = await self.cross_rag.search(query, top_k, layer_list)
return [
{
"chunk_id": r.chunk_id,
"title": r.content.split('\n')[0][:50] if r.content else r.chunk_id,
"content": r.content[:200] + "..." if len(r.content) > 200 else r.content,
"score": r.score,
"layer": r.source_layer.value,
"project": r.source_project,
"chapter": r.chapter,
"type": r.chunk_type
}
for r in results
]
# ==================== 添加 ====================
def add_pattern(
self,
pattern_id: str,
pattern_type: str,
title: str,
description: str,
catharsis_model: str,
source_chapter: int,
layer: str = "project",
tension_curve: Optional[list] = None,
structure: Optional[dict] = None,
style_tags: Optional[list] = None
) -> bool:
"""添加模式"""
pattern = LearnedPattern(
pattern_id=pattern_id,
pattern_type=pattern_type,
title=title,
description=description,
tension_curve=tension_curve or [],
catharsis_model=catharsis_model,
structure=structure or {},
hot_spots=[],
style_tags=style_tags or [],
source_project=self.project_root.name,
source_chapter=source_chapter,
learned_at=datetime.now().isoformat()
)
rag_layer = RAGLayer.SYSTEM if layer == "system" else RAGLayer.PROJECT
return self.cross_rag.store_learned_pattern(pattern, rag_layer)
def learn_and_add(
self,
chapter_file: Path,
layer: str = "project"
) -> Optional[str]:
"""从章节学习并添加模式"""
# 分析章节
result = self.analyzer.analyze_chapter(chapter_file)
# 生成模式
pattern = self.analyzer.learn_pattern(result, layer)
# 存储
rag_layer = RAGLayer.SYSTEM if layer == "system" else RAGLayer.PROJECT
success = self.cross_rag.store_learned_pattern(pattern, rag_layer)
if success:
return pattern.pattern_id
return None
# ==================== 删除 ====================
def delete_pattern(self, pattern_id: str, layer: str) -> bool:
"""删除模式"""
if layer == "project":
return self._delete_from_project_db(pattern_id)
elif layer == "system":
return self._delete_from_system(pattern_id)
return False
def _delete_from_project_db(self, pattern_id: str) -> bool:
"""从项目数据库删除"""
db_path = self.project_root / ".noma" / "rag" / "learned.db"
if not db_path.exists():
return False
try:
import sqlite3
conn = sqlite3.connect(str(db_path))
cursor = conn.cursor()
cursor.execute("DELETE FROM learned_patterns WHERE pattern_id = ?", (pattern_id,))
affected = cursor.rowcount
conn.commit()
conn.close()
return affected > 0
except Exception:
return False
def _delete_from_system(self, pattern_id: str) -> bool:
"""从系统目录删除"""
pattern_file = self.cross_rag.system_learned_dir / f"{pattern_id}.json"
if pattern_file.exists():
pattern_file.unlink()
return True
return False
# ==================== 列表 ====================
def list_patterns(self, layer: str) -> list:
"""列出指定层的模式"""
if layer == "project":
return self._list_project_patterns()
elif layer == "system":
return self._list_system_patterns()
elif layer == "plugin":
return self._list_plugin_patterns()
return []
def _list_project_patterns(self) -> list:
"""列出项目层模式"""
db_path = self.project_root / ".noma" / "rag" / "learned.db"
if not db_path.exists():
return []
try:
import sqlite3
conn = sqlite3.connect(str(db_path))
cursor = conn.cursor()
cursor.execute("""
SELECT pattern_id, pattern_type, title, description,
catharsis_model, source_chapter, learned_at, usage_count
FROM learned_patterns
ORDER BY learned_at DESC
""")
rows = cursor.fetchall()
conn.close()
return [
{
"pattern_id": r[0],
"type": r[1],
"title": r[2],
"description": r[3][:100] + "..." if r[3] and len(r[3]) > 100 else r[3],
"catharsis_model": r[4],
"source_chapter": r[5],
"learned_at": r[6],
"usage_count": r[7]
}
for r in rows
]
except Exception:
return []
def _list_system_patterns(self) -> list:
"""列出系统层模式"""
if not self.cross_rag.system_learned_dir.exists():
return []
patterns = []
for f in self.cross_rag.system_learned_dir.glob("*.json"):
try:
data = json.loads(f.read_text(encoding="utf-8"))
patterns.append({
"pattern_id": data.get("pattern_id", f.stem),
"type": data.get("pattern_type", "unknown"),
"title": data.get("title", f.stem),
"description": data.get("description", "")[:100],
"catharsis_model": data.get("catharsis_model", "unknown"),
"source_project": data.get("source_project", "unknown"),
"source_chapter": data.get("source_chapter", 0),
"learned_at": data.get("learned_at", "")
})
except Exception:
continue
return sorted(patterns, key=lambda x: x.get("learned_at", ""), reverse=True)
def _list_plugin_patterns(self) -> list:
"""列出插件层模式"""
patterns = []
plugin_dir = self.cross_rag.plugin_matrices_dir
# catharsis models
catharsis_dir = plugin_dir / "catharsis_models"
if catharsis_dir.exists():
for f in catharsis_dir.glob("*.md"):
patterns.append({
"pattern_id": f"plugin:{f.stem}",
"type": "catharsis_model",
"title": f.stem,
"description": "内置爽感模型",
"catharsis_model": f.stem,
"source": "noma_plugin"
})
return patterns
# ==================== 同步 ====================
def sync_to_system(self, pattern_id: str) -> bool:
"""将项目模式同步到系统层"""
# 从项目数据库读取
db_path = self.project_root / ".noma" / "rag" / "learned.db"
if not db_path.exists():
return False
try:
import sqlite3
conn = sqlite3.connect(str(db_path))
cursor = conn.cursor()
cursor.execute("""
SELECT pattern_id, pattern_type, title, description,
tension_curve, catharsis_model, structure,
hot_spots, style_tags, source_project,
source_chapter, learned_at, usage_count
FROM learned_patterns WHERE pattern_id = ?
""", (pattern_id,))
row = cursor.fetchone()
conn.close()
if not row:
return False
pattern = LearnedPattern(
pattern_id=row[0],
pattern_type=row[1],
title=row[2],
description=row[3] or "",
tension_curve=json.loads(row[4]) if row[4] else [],
catharsis_model=row[5] or "unknown",
structure=json.loads(row[6]) if row[6] else {},
hot_spots=json.loads(row[7]) if row[7] else [],
style_tags=json.loads(row[8]) if row[8] else [],
source_project=row[9] or self.project_root.name,
source_chapter=row[10] or 0,
learned_at=row[11] or datetime.now().isoformat(),
usage_count=row[12] or 0
)
return self.cross_rag.store_learned_pattern(pattern, RAGLayer.SYSTEM)
except Exception:
return False
def sync_all_to_system(self) -> dict:
"""同步所有项目模式到系统层"""
project_patterns = self._list_project_patterns()
synced = 0
failed = 0
for p in project_patterns:
if self.sync_to_system(p["pattern_id"]):
synced += 1
else:
failed += 1
return {"synced": synced, "failed": failed, "total": len(project_patterns)}
# ==================== CLI ====================
def main():
if sys.platform == "win32":
enable_windows_utf8_stdio()
parser = argparse.ArgumentParser(description="RAG Manager - RAG 检索与管理")
parser.add_argument("--project-root", type=str, default=".",
help="项目根目录")
subparsers = parser.add_subparsers(dest="command")
# 搜索
search_parser = subparsers.add_parser("search", help="检索 RAG")
search_parser.add_argument("query", help="检索 query")
search_parser.add_argument("--top-k", type=int, default=5)
search_parser.add_argument("--layers", type=str, default="project,system",
help="检索层级,逗号分隔")
# 列表
list_parser = subparsers.add_parser("list", help="列出模式")
list_parser.add_argument("--layer", choices=["project", "system", "plugin"],
default="project", help="RAG 层")
# 添加
add_parser = subparsers.add_parser("add", help="添加模式")
add_parser.add_argument("--pattern-id", required=True)
add_parser.add_argument("--pattern-type", required=True)
add_parser.add_argument("--title", required=True)
add_parser.add_argument("--description", required=True)
add_parser.add_argument("--catharsis-model", required=True)
add_parser.add_argument("--source-chapter", type=int, required=True)
add_parser.add_argument("--layer", choices=["project", "system"],
default="project")
add_parser.add_argument("--learn", help="从章节文件学习")
add_parser.add_argument("--chapter-file", help="章节文件路径")
# 删除
delete_parser = subparsers.add_parser("delete", help="删除模式")
delete_parser.add_argument("--pattern-id", required=True)
delete_parser.add_argument("--layer", choices=["project", "system"],
required=True)
# 同步
sync_parser = subparsers.add_parser("sync", help="同步到系统层")
sync_parser.add_argument("--pattern-id", help="同步单个模式(可选)")
sync_parser.add_argument("--all", action="store_true", help="同步所有")
args = parser.parse_args()
if not args.project_root:
print("Error: --project-root is required")
sys.exit(1)
project_root = Path(args.project_root).resolve()
manager = RAGManager(project_root)
if args.command == "search":
layers = [l.strip() for l in args.layers.split(",")]
results = asyncio.run(manager.search(args.query, args.top_k, layers))
print(f"\n=== Search Results ({len(results)}) ===")
for r in results:
print(f"\n[{r['layer']}] {r['title']}")
print(f" Type: {r['type']}, Chapter: {r['chapter']}")
print(f" Score: {r['score']:.2f}")
print(f" Content: {r['content']}")
elif args.command == "list":
patterns = manager.list_patterns(args.layer)
print(f"\n=== {args.layer.upper()} Patterns ({len(patterns)}) ===")
for p in patterns:
print(f"\n[{p['pattern_id']}] {p['title']}")
print(f" Type: {p['type']}, Catharsis: {p.get('catharsis_model', 'N/A')}")
print(f" Source: {p.get('source_project', 'N/A')} Ch.{p.get('source_chapter', 0)}")
if p.get("description"):
print(f" Desc: {p['description'][:100]}")
elif args.command == "add":
if hasattr(args, 'learn') and args.learn:
# 从章节学习
chapter_file = project_root / args.chapter_file
pattern_id = manager.learn_and_add(chapter_file, args.layer)
if pattern_id:
print(f"✓ Pattern learned and added: {pattern_id}")
else:
print("✗ Failed to learn pattern")
sys.exit(1)
else:
success = manager.add_pattern(
args.pattern_id,
args.pattern_type,
args.title,
args.description,
args.catharsis_model,
args.source_chapter,
args.layer
)
if success:
print(f"✓ Pattern added to {args.layer}")
else:
print("✗ Failed to add pattern")
sys.exit(1)
elif args.command == "delete":
success = manager.delete_pattern(args.pattern_id, args.layer)
if success:
print(f"✓ Pattern deleted from {args.layer}")
else:
print("✗ Failed to delete pattern")
sys.exit(1)
elif args.command == "sync":
if args.all:
result = manager.sync_all_to_system()
print(f"✓ Synced {result['synced']}/{result['total']} patterns")
if result['failed'] > 0:
print(f" Failed: {result['failed']}")
elif args.pattern_id:
success = manager.sync_to_system(args.pattern_id)
if success:
print(f"✓ Pattern synced to system")
else:
print("✗ Failed to sync pattern")
sys.exit(1)
else:
print("Specify --pattern-id or --all")
sys.exit(1)
else:
parser.print_help()
if __name__ == "__main__":
main()
+125
View File
@@ -0,0 +1,125 @@
#!/usr/bin/env python3
# -*- coding: utf-8 -*-
"""
Pydantic schemas for data_modules outputs.
"""
from __future__ import annotations
from typing import Any, Dict, List, Optional
from pydantic import BaseModel, Field, ValidationError, ConfigDict
class EntityAppeared(BaseModel):
model_config = ConfigDict(extra="allow")
id: str
type: str
mentions: List[str] = Field(default_factory=list)
confidence: float = 1.0
class EntityNew(BaseModel):
model_config = ConfigDict(extra="allow")
suggested_id: str
name: str
type: str
tier: str = "装饰"
class StateChange(BaseModel):
model_config = ConfigDict(extra="allow")
entity_id: str
field: str
old: Optional[str] = None
new: str
reason: Optional[str] = None
class RelationshipNew(BaseModel):
model_config = ConfigDict(extra="allow", populate_by_name=True)
from_entity: str = Field(alias="from")
to_entity: str = Field(alias="to")
type: str
description: Optional[str] = None
chapter: Optional[int] = None
class UncertainCandidate(BaseModel):
model_config = ConfigDict(extra="allow")
type: str
id: str
class UncertainMention(BaseModel):
model_config = ConfigDict(extra="allow")
mention: str
candidates: List[UncertainCandidate] = Field(default_factory=list)
confidence: float = 0.0
adopted: Optional[str] = None
class DataAgentOutput(BaseModel):
model_config = ConfigDict(extra="allow")
entities_appeared: List[EntityAppeared] = Field(default_factory=list)
entities_new: List[EntityNew] = Field(default_factory=list)
state_changes: List[StateChange] = Field(default_factory=list)
relationships_new: List[RelationshipNew] = Field(default_factory=list)
scenes_chunked: int = 0
uncertain: List[UncertainMention] = Field(default_factory=list)
warnings: List[str] = Field(default_factory=list)
class ErrorSchema(BaseModel):
model_config = ConfigDict(extra="allow")
code: str
message: str
suggestion: Optional[str] = None
details: Optional[Dict[str, Any]] = None
def validate_data_agent_output(payload: Dict[str, Any]) -> DataAgentOutput:
return DataAgentOutput.model_validate(payload)
def format_validation_error(exc: ValidationError) -> Dict[str, Any]:
return {
"code": "SCHEMA_VALIDATION_FAILED",
"message": "数据结构校验失败",
"details": {"errors": exc.errors()},
"suggestion": "请检查 data-agent 输出字段是否完整且类型正确",
}
def normalize_data_agent_output(payload: Dict[str, Any]) -> Dict[str, Any]:
if not isinstance(payload, dict):
return {}
def _ensure_list(key: str):
value = payload.get(key)
if value is None:
payload[key] = []
elif isinstance(value, list):
return
else:
payload[key] = [value]
for key in [
"entities_appeared",
"entities_new",
"state_changes",
"relationships_new",
"uncertain",
"warnings",
]:
_ensure_list(key)
payload.setdefault("scenes_chunked", 0)
return payload
@@ -0,0 +1,92 @@
#!/usr/bin/env python3
# -*- coding: utf-8 -*-
"""
Context snapshot manager.
"""
from __future__ import annotations
import json
from dataclasses import dataclass
from datetime import datetime, timezone
from filelock import FileLock
from pathlib import Path
from typing import Any, Dict, Optional
from .config import get_config
try:
# 当 scripts 目录在 sys.path 中
from security_utils import atomic_write_json
except ImportError: # pragma: no cover
# 当以 python -m scripts.data_modules... 形式运行
from scripts.security_utils import atomic_write_json
SNAPSHOT_VERSION = "1.2"
class SnapshotVersionMismatch(RuntimeError):
def __init__(self, expected: str, actual: str) -> None:
super().__init__(f"snapshot version mismatch: expected {expected}, got {actual}")
self.expected = expected
self.actual = actual
@dataclass
class SnapshotMeta:
chapter: int
version: str
saved_at: str
class SnapshotManager:
def __init__(self, config=None, version: str = SNAPSHOT_VERSION):
self.config = config or get_config()
self.version = version
self.snapshot_dir = self.config.noma_dir / "context_snapshots"
self.snapshot_dir.mkdir(parents=True, exist_ok=True)
def _snapshot_path(self, chapter: int) -> Path:
return self.snapshot_dir / f"ch{chapter:04d}.json"
def _snapshot_lock_path(self, chapter: int) -> Path:
return self._snapshot_path(chapter).with_suffix(".json.lock")
def save_snapshot(self, chapter: int, payload: Dict[str, Any], meta: Optional[Dict[str, Any]] = None) -> Path:
data: Dict[str, Any] = {
"version": self.version,
"chapter": chapter,
"saved_at": datetime.now(timezone.utc).isoformat(),
"payload": payload,
}
if meta:
data["meta"] = meta
path = self._snapshot_path(chapter)
lock = FileLock(str(self._snapshot_lock_path(chapter)), timeout=10)
with lock:
atomic_write_json(path, data, use_lock=False, backup=False)
return path
def load_snapshot(self, chapter: int) -> Optional[Dict[str, Any]]:
path = self._snapshot_path(chapter)
lock = FileLock(str(self._snapshot_lock_path(chapter)), timeout=10)
with lock:
if not path.exists():
return None
data = json.loads(path.read_text(encoding="utf-8"))
version = str(data.get("version", ""))
if version != self.version:
raise SnapshotVersionMismatch(self.version, version)
return data
def delete_snapshot(self, chapter: int) -> bool:
path = self._snapshot_path(chapter)
lock = FileLock(str(self._snapshot_lock_path(chapter)), timeout=10)
with lock:
if path.exists():
path.unlink()
return True
return False
def list_snapshots(self) -> list[str]:
return sorted(p.name for p in self.snapshot_dir.glob("ch*.json"))
@@ -0,0 +1,594 @@
#!/usr/bin/env python3
# -*- coding: utf-8 -*-
"""
SQL State Manager - SQLite 状态管理模块 (v5.4)
基于 IndexManager 扩展,提供与 StateManager 兼容的高级接口,
将大数据(实体、别名、状态变化、关系)存储到 SQLite 而非 JSON。
目标(v5.1 引入,v5.4 沿用):
- 替代 state.json 中的大数据字段
- 保持与 Data Agent / Context Agent 的接口兼容
- 支持增量写入和按需查询
"""
import json
from typing import Dict, List, Optional, Any
from dataclasses import dataclass, field
from datetime import datetime
from .index_manager import (
IndexManager,
EntityMeta,
StateChangeMeta,
RelationshipMeta,
RelationshipEventMeta,
)
from .config import get_config
from .observability import safe_log_tool_call
@dataclass
class EntityData:
"""实体数据(用于 Data Agent 输入)"""
id: str
type: str # 角色/地点/物品/势力/招式
name: str
tier: str = "装饰"
desc: str = ""
current: Dict[str, Any] = field(default_factory=dict)
aliases: List[str] = field(default_factory=list)
first_appearance: int = 0
last_appearance: int = 0
is_protagonist: bool = False
class SQLStateManager:
"""
SQLite 状态管理器(v5.1 引入,v5.4 沿用)
提供与 StateManager 兼容的接口,但数据存储在 SQLite (index.db) 中。
用于替代 state.json 中膨胀的数据结构。
用法:
```python
manager = SQLStateManager(config)
# 写入实体
manager.upsert_entity(EntityData(
id="xiaoyan",
type="角色",
name="萧炎",
tier="核心",
current={"realm": "斗师", "location": "天云宗"},
aliases=["小炎子", "废柴"],
is_protagonist=True
))
# 写入状态变化
manager.record_state_change(
entity_id="xiaoyan",
field="realm",
old_value="斗者",
new_value="斗师",
reason="闭关突破",
chapter=100
)
# 写入关系
manager.upsert_relationship(
from_entity="xiaoyan",
to_entity="yaolao",
type="师徒",
description="药老收萧炎为徒",
chapter=5
)
# 读取
protagonist = manager.get_protagonist()
core_entities = manager.get_core_entities()
changes = manager.get_recent_state_changes(limit=50)
```
"""
# v5.0 引入的实体类型
ENTITY_TYPES = ["角色", "地点", "物品", "势力", "招式"]
def __init__(self, config=None):
self.config = config or get_config()
self._index_manager = IndexManager(config)
# ==================== 实体操作 ====================
def upsert_entity(self, entity: EntityData) -> bool:
"""
插入或更新实体
自动处理:
- 实体基本信息写入 entities 表
- 别名写入 aliases 表
- canonical_name 自动添加为别名
返回: 是否为新实体
"""
# 构建 EntityMeta
meta = EntityMeta(
id=entity.id,
type=entity.type,
canonical_name=entity.name,
tier=entity.tier,
desc=entity.desc,
current=entity.current,
first_appearance=entity.first_appearance,
last_appearance=entity.last_appearance,
is_protagonist=entity.is_protagonist,
is_archived=False
)
is_new = self._index_manager.upsert_entity(meta)
# 注册别名
# 1. canonical_name 本身作为别名
self._index_manager.register_alias(entity.name, entity.id, entity.type)
# 2. 其他别名
for alias in entity.aliases:
if alias and alias != entity.name:
self._index_manager.register_alias(alias, entity.id, entity.type)
return is_new
def get_entity(self, entity_id: str) -> Optional[Dict]:
"""获取实体详情"""
entity = self._index_manager.get_entity(entity_id)
if entity:
# 添加别名
entity["aliases"] = self._index_manager.get_entity_aliases(entity_id)
return entity
def get_entities_by_type(self, entity_type: str, include_archived: bool = False) -> List[Dict]:
"""按类型获取实体"""
entities = self._index_manager.get_entities_by_type(entity_type, include_archived)
for e in entities:
e["aliases"] = self._index_manager.get_entity_aliases(e["id"])
return entities
def get_core_entities(self) -> List[Dict]:
"""
获取核心实体(用于 Context Agent 全量加载)
返回所有 tier=核心/重要 或 is_protagonist=1 的实体
(次要/装饰实体按需查询,不全量加载)
"""
entities = self._index_manager.get_core_entities()
for e in entities:
e["aliases"] = self._index_manager.get_entity_aliases(e["id"])
return entities
def get_protagonist(self) -> Optional[Dict]:
"""获取主角实体"""
protagonist = self._index_manager.get_protagonist()
if protagonist:
protagonist["aliases"] = self._index_manager.get_entity_aliases(protagonist["id"])
return protagonist
def update_entity_current(self, entity_id: str, updates: Dict) -> bool:
"""增量更新实体的 current 字段"""
return self._index_manager.update_entity_current(entity_id, updates)
def resolve_alias(self, alias: str) -> List[Dict]:
"""
根据别名解析实体(一对多)
返回所有匹配的实体
"""
return self._index_manager.get_entities_by_alias(alias)
def register_alias(self, alias: str, entity_id: str, entity_type: str) -> bool:
"""注册别名"""
return self._index_manager.register_alias(alias, entity_id, entity_type)
# ==================== 状态变化操作 ====================
def record_state_change(
self,
entity_id: str,
field: str,
old_value: Any,
new_value: Any,
reason: str,
chapter: int
) -> int:
"""
记录状态变化
返回: 记录 ID
"""
change = StateChangeMeta(
entity_id=entity_id,
field=field,
old_value=str(old_value) if old_value is not None else "",
new_value=str(new_value),
reason=reason,
chapter=chapter
)
return self._index_manager.record_state_change(change)
def get_entity_state_changes(self, entity_id: str, limit: int = 20) -> List[Dict]:
"""获取实体的状态变化历史"""
return self._index_manager.get_entity_state_changes(entity_id, limit)
def get_recent_state_changes(self, limit: int = 50) -> List[Dict]:
"""获取最近的状态变化"""
return self._index_manager.get_recent_state_changes(limit)
def get_chapter_state_changes(self, chapter: int) -> List[Dict]:
"""获取某章的所有状态变化"""
return self._index_manager.get_chapter_state_changes(chapter)
# ==================== 关系操作 ====================
def upsert_relationship(
self,
from_entity: str,
to_entity: str,
type: str,
description: str,
chapter: int
) -> bool:
"""
插入或更新关系
返回: 是否为新关系
"""
rel = RelationshipMeta(
from_entity=from_entity,
to_entity=to_entity,
type=type,
description=description,
chapter=chapter
)
return self._index_manager.upsert_relationship(rel)
def get_entity_relationships(self, entity_id: str, direction: str = "both") -> List[Dict]:
"""获取实体的关系"""
return self._index_manager.get_entity_relationships(entity_id, direction)
def get_relationship_between(self, entity1: str, entity2: str) -> List[Dict]:
"""获取两个实体之间的所有关系"""
return self._index_manager.get_relationship_between(entity1, entity2)
def get_recent_relationships(self, limit: int = 30) -> List[Dict]:
"""获取最近建立的关系"""
return self._index_manager.get_recent_relationships(limit)
# ==================== 批量写入(供 Data Agent 使用) ====================
def process_chapter_entities(
self,
chapter: int,
entities_appeared: List[Dict],
entities_new: List[Dict],
state_changes: List[Dict],
relationships_new: List[Dict]
) -> Dict[str, int]:
"""
处理章节的实体数据(Data Agent 主入口)
参数:
- chapter: 章节号
- entities_appeared: 出场的已有实体
[{"id": "xiaoyan", "type": "角色", "mentions": ["萧炎", "他"], "confidence": 0.95}]
- entities_new: 新发现的实体
[{"suggested_id": "hongyi_girl", "name": "红衣女子", "type": "角色", "tier": "装饰"}]
- state_changes: 状态变化
[{"entity_id": "xiaoyan", "field": "realm", "old": "斗者", "new": "斗师", "reason": "突破"}]
- relationships_new: 新关系
[{"from": "xiaoyan", "to": "hongyi_girl", "type": "相识", "description": "初次见面"}]
返回: 写入统计
"""
stats = {
"entities_updated": 0,
"entities_created": 0,
"state_changes": 0,
"relationships": 0,
"aliases": 0
}
# 1. 处理出场实体(更新 last_appearance)
for entity in entities_appeared:
entity_id = entity.get("id")
if not entity_id:
continue
self._index_manager.update_entity_current(entity_id, {}) # 触发 updated_at
# 更新 last_appearance
existing = self._index_manager.get_entity(entity_id)
if existing:
# 使用 SQL 直接更新 last_appearance
self._update_last_appearance(entity_id, chapter)
stats["entities_updated"] += 1
# 记录出场(保留原有逻辑)
self._index_manager.record_appearance(
entity_id=entity_id,
chapter=chapter,
mentions=entity.get("mentions", []),
confidence=entity.get("confidence", 1.0)
)
# 2. 处理新实体
for entity in entities_new:
suggested_id = entity.get("suggested_id") or entity.get("id")
if not suggested_id:
continue
entity_data = EntityData(
id=suggested_id,
type=entity.get("type", "角色"),
name=entity.get("name", suggested_id),
tier=entity.get("tier", "装饰"),
desc=entity.get("desc", ""),
current=entity.get("current", {}),
aliases=entity.get("aliases", []),
first_appearance=chapter,
last_appearance=chapter,
is_protagonist=entity.get("is_protagonist", False)
)
is_new = self.upsert_entity(entity_data)
if is_new:
stats["entities_created"] += 1
else:
stats["entities_updated"] += 1
# 统计别名
stats["aliases"] += 1 + len(entity_data.aliases)
# 记录新实体的首次出场(解决 appearances 缺失问题)
mentions = entity.get("mentions", [])
if not mentions:
mentions = [entity_data.name] # 至少包含实体名
self._index_manager.record_appearance(
entity_id=suggested_id,
chapter=chapter,
mentions=mentions,
confidence=entity.get("confidence", 1.0)
)
# 3. 处理状态变化
for change in state_changes:
entity_id = change.get("entity_id")
if not entity_id:
continue
self.record_state_change(
entity_id=entity_id,
field=change.get("field", ""),
old_value=change.get("old", change.get("old_value", "")),
new_value=change.get("new", change.get("new_value", "")),
reason=change.get("reason", ""),
chapter=chapter
)
stats["state_changes"] += 1
# 同步更新实体的 current
field_name = change.get("field")
new_value = change.get("new", change.get("new_value"))
# 注意:new_value 可能是 0/""/False 等 falsy 值,需要用 is not None 判断
if field_name and new_value is not None:
self._index_manager.update_entity_current(entity_id, {field_name: new_value})
# 4. 处理新关系
for rel in relationships_new:
from_entity = rel.get("from", rel.get("from_entity"))
to_entity = rel.get("to", rel.get("to_entity"))
if not from_entity or not to_entity:
continue
rel_type = rel.get("type", "相识")
description = rel.get("description", "")
# v5.5: 先记录关系事件,再更新关系快照
self._index_manager.record_relationship_event(
RelationshipEventMeta(
from_entity=from_entity,
to_entity=to_entity,
type=rel_type,
chapter=chapter,
action=rel.get("action", "update"),
polarity=rel.get("polarity", 0),
strength=rel.get("strength", 0.5),
description=description,
scene_index=rel.get("scene_index", 0),
evidence=rel.get("evidence", ""),
confidence=rel.get("confidence", 1.0),
)
)
self.upsert_relationship(
from_entity=from_entity,
to_entity=to_entity,
type=rel_type,
description=description,
chapter=chapter
)
stats["relationships"] += 1
return stats
def _update_last_appearance(self, entity_id: str, chapter: int):
"""更新实体的 last_appearance"""
with self._index_manager._get_conn() as conn:
cursor = conn.cursor()
cursor.execute("""
UPDATE entities SET
last_appearance = MAX(last_appearance, ?),
updated_at = CURRENT_TIMESTAMP
WHERE id = ?
""", (chapter, entity_id))
conn.commit()
# ==================== 统计 ====================
def get_stats(self) -> Dict[str, int]:
"""获取统计信息"""
return self._index_manager.get_stats()
# ==================== 格式转换(兼容性) ====================
def export_to_entities_v3_format(self) -> Dict[str, Dict[str, Dict]]:
"""
导出为 entities_v3 格式(用于兼容性)
返回: {"角色": {"xiaoyan": {...}}, "地点": {...}, ...}
"""
result = {t: {} for t in self.ENTITY_TYPES}
for entity_type in self.ENTITY_TYPES:
entities = self.get_entities_by_type(entity_type, include_archived=True)
for e in entities:
entity_dict = {
"canonical_name": e.get("canonical_name"),
"name": e.get("canonical_name"), # 兼容性别名
"tier": e.get("tier", "装饰"),
"aliases": e.get("aliases", []),
"desc": e.get("desc", ""),
"current": e.get("current_json", {}),
"history": [], # 历史记录需要从 state_changes 表查询
"first_appearance": e.get("first_appearance", 0),
"last_appearance": e.get("last_appearance", 0)
}
if e.get("is_protagonist"):
entity_dict["is_protagonist"] = True
result[entity_type][e["id"]] = entity_dict
return result
def export_to_alias_index_format(self) -> Dict[str, List[Dict[str, str]]]:
"""
导出为 alias_index 格式(用于兼容性)
返回: {"萧炎": [{"type": "角色", "id": "xiaoyan"}], ...}
"""
result = {}
with self._index_manager._get_conn() as conn:
cursor = conn.cursor()
cursor.execute("SELECT alias, entity_id, entity_type FROM aliases")
for row in cursor.fetchall():
alias = row["alias"]
if alias not in result:
result[alias] = []
result[alias].append({
"type": row["entity_type"],
"id": row["entity_id"]
})
return result
# ==================== CLI 接口 ====================
def main():
import argparse
import sys
from .cli_output import print_success, print_error
from .cli_args import normalize_global_project_root, load_json_arg
from .index_manager import IndexManager
parser = argparse.ArgumentParser(description="SQL State Manager CLI (v5.4)")
parser.add_argument("--project-root", type=str, help="项目根目录")
subparsers = parser.add_subparsers(dest="command")
# 获取统计
subparsers.add_parser("stats")
# 获取主角
subparsers.add_parser("get-protagonist")
# 获取核心实体
subparsers.add_parser("get-core-entities")
# 导出 entities_v3 格式
subparsers.add_parser("export-entities-v3")
# 导出 alias_index 格式
subparsers.add_parser("export-alias-index")
# 处理章节数据
process_parser = subparsers.add_parser("process-chapter")
process_parser.add_argument("--chapter", type=int, required=True)
process_parser.add_argument("--data", required=True, help="JSON 格式的章节数据")
argv = normalize_global_project_root(sys.argv[1:])
args = parser.parse_args(argv)
# 初始化
config = None
if args.project_root:
# 允许传入“工作区根目录”,统一解析到真正的 book project_root(必须包含 .noma/state.json)
from project_locator import resolve_project_root
from .config import DataModulesConfig
resolved_root = resolve_project_root(args.project_root)
config = DataModulesConfig.from_project_root(resolved_root)
manager = SQLStateManager(config)
logger = IndexManager(config)
tool_name = f"sql_state_manager:{args.command or 'unknown'}"
def emit_success(data=None, message: str = "ok"):
print_success(data, message=message)
safe_log_tool_call(logger, tool_name=tool_name, success=True)
def emit_error(code: str, message: str, suggestion: str | None = None):
print_error(code, message, suggestion=suggestion)
safe_log_tool_call(
logger,
tool_name=tool_name,
success=False,
error_code=code,
error_message=message,
)
if args.command == "stats":
stats = manager.get_stats()
emit_success(stats, message="stats")
elif args.command == "get-protagonist":
protagonist = manager.get_protagonist()
if protagonist:
emit_success(protagonist, message="protagonist")
else:
emit_error("NOT_FOUND", "未设置主角")
elif args.command == "get-core-entities":
entities = manager.get_core_entities()
emit_success(entities, message="core_entities")
elif args.command == "export-entities-v3":
data = manager.export_to_entities_v3_format()
emit_success(data, message="entities_v3")
elif args.command == "export-alias-index":
data = manager.export_to_alias_index_format()
emit_success(data, message="alias_index")
elif args.command == "process-chapter":
data = load_json_arg(args.data)
stats = manager.process_chapter_entities(
chapter=args.chapter,
entities_appeared=data.get("entities_appeared", []),
entities_new=data.get("entities_new", []),
state_changes=data.get("state_changes", []),
relationships_new=data.get("relationships_new", []),
)
emit_success(stats, message="chapter_processed")
else:
emit_error("UNKNOWN_COMMAND", "未指定有效命令", suggestion="请查看 --help")
if __name__ == "__main__":
main()
File diff suppressed because it is too large. Load diff
@@ -0,0 +1,249 @@
#!/usr/bin/env python3
# -*- coding: utf-8 -*-
"""
Runtime validators/normalizers for state.json sections.
"""
from __future__ import annotations
import re
from typing import Any, Dict, List, Mapping, Optional, Sequence
FORESHADOWING_STATUS_PENDING = "未回收"
FORESHADOWING_STATUS_RESOLVED = "已回收"
FORESHADOWING_TIER_CORE = "核心"
FORESHADOWING_TIER_SUB = "支线"
FORESHADOWING_TIER_DECOR = "装饰"
FORESHADOWING_PLANTED_KEYS = [
"planted_chapter",
"added_chapter",
"source_chapter",
"start_chapter",
"chapter",
]
FORESHADOWING_TARGET_KEYS = [
"target_chapter",
"due_chapter",
"deadline_chapter",
"resolve_by_chapter",
"target",
]
_PENDING_STATUS_TEXT = {"未回收", "待回收", "进行中", "未解决", "pending", "active"}
_RESOLVED_STATUS_TEXT = {"已回收", "已完成", "已解决", "完成", "resolved", "done", "complete"}
_TIER_CORE_TEXT = {"核心", "主线", "core", "main"}
_TIER_DECOR_TEXT = {"装饰", "次要", "decor", "decoration"}
_PATTERN_FIELDS = [
"coolpoint_patterns",
"coolpoint_pattern",
"cool_point_patterns",
"cool_point_pattern",
"patterns",
"pattern",
]
_PATTERN_SPLIT_RE = re.compile(r"[、,,/|+;;。]+")
def to_positive_int(value: Any) -> Optional[int]:
if value is None or isinstance(value, bool):
return None
try:
number = int(value)
return number if number > 0 else None
except (TypeError, ValueError):
if isinstance(value, str):
matched = re.search(r"\d+", value)
if matched:
number = int(matched.group(0))
return number if number > 0 else None
return None
def resolve_chapter_field(item: Mapping[str, Any], keys: Sequence[str]) -> Optional[int]:
for key in keys:
if key in item:
chapter = to_positive_int(item.get(key))
if chapter is not None:
return chapter
return None
def normalize_foreshadowing_status(
raw_status: Any,
default: str = FORESHADOWING_STATUS_PENDING,
) -> str:
text = str(raw_status or "").strip()
if not text:
return default
text_lower = text.lower()
if (
text in _RESOLVED_STATUS_TEXT
or text_lower in _RESOLVED_STATUS_TEXT
or FORESHADOWING_STATUS_RESOLVED in text
):
return FORESHADOWING_STATUS_RESOLVED
if text in _PENDING_STATUS_TEXT or text_lower in _PENDING_STATUS_TEXT:
return FORESHADOWING_STATUS_PENDING
return default
def is_resolved_foreshadowing_status(raw_status: Any) -> bool:
return normalize_foreshadowing_status(raw_status) == FORESHADOWING_STATUS_RESOLVED
def normalize_foreshadowing_tier(
raw_tier: Any,
default: str = FORESHADOWING_TIER_SUB,
) -> str:
text = str(raw_tier or "").strip()
if not text:
return default
text_lower = text.lower()
if text in _TIER_CORE_TEXT or text_lower in _TIER_CORE_TEXT:
return FORESHADOWING_TIER_CORE
if text in _TIER_DECOR_TEXT or text_lower in _TIER_DECOR_TEXT:
return FORESHADOWING_TIER_DECOR
return default
def split_patterns(raw_value: Any) -> List[str]:
if raw_value is None:
return []
tokens: List[str] = []
if isinstance(raw_value, list):
for item in raw_value:
text = str(item).strip()
if text:
tokens.append(text)
elif isinstance(raw_value, str):
text = raw_value.strip()
if not text:
return []
split_values = [part.strip() for part in _PATTERN_SPLIT_RE.split(text)]
tokens.extend([part for part in split_values if part])
else:
return []
deduped: List[str] = []
seen = set()
for token in tokens:
if token not in seen:
seen.add(token)
deduped.append(token)
return deduped
def count_patterns(raw_value: Any) -> Optional[int]:
patterns = split_patterns(raw_value)
if not patterns:
return None
return len(patterns)
def normalize_foreshadowing_item(item: Mapping[str, Any]) -> Dict[str, Any]:
normalized = dict(item)
normalized["status"] = normalize_foreshadowing_status(item.get("status"))
normalized["tier"] = normalize_foreshadowing_tier(item.get("tier"))
content = str(item.get("content") or "").strip()
if content:
normalized["content"] = content
planted_chapter = resolve_chapter_field(item, FORESHADOWING_PLANTED_KEYS)
if planted_chapter is not None:
normalized["planted_chapter"] = planted_chapter
target_chapter = resolve_chapter_field(item, FORESHADOWING_TARGET_KEYS)
if target_chapter is not None:
normalized["target_chapter"] = target_chapter
resolved_chapter = resolve_chapter_field(item, ["resolved_chapter", "resolved_at_chapter", "resolved"])
if resolved_chapter is not None:
normalized["resolved_chapter"] = resolved_chapter
return normalized
def normalize_foreshadowing_list(raw_items: Any) -> List[Dict[str, Any]]:
if not isinstance(raw_items, list):
return []
normalized: List[Dict[str, Any]] = []
for raw_item in raw_items:
if isinstance(raw_item, Mapping):
normalized.append(normalize_foreshadowing_item(raw_item))
return normalized
def normalize_chapter_meta_entry(entry: Mapping[str, Any]) -> Dict[str, Any]:
normalized = dict(entry)
merged_patterns: List[str] = []
seen = set()
for field_name in _PATTERN_FIELDS:
for pattern in split_patterns(entry.get(field_name)):
if pattern not in seen:
seen.add(pattern)
merged_patterns.append(pattern)
if merged_patterns:
normalized["coolpoint_patterns"] = merged_patterns
return normalized
def normalize_chapter_meta(raw_chapter_meta: Any) -> Dict[str, Dict[str, Any]]:
if not isinstance(raw_chapter_meta, Mapping):
return {}
normalized: Dict[str, Dict[str, Any]] = {}
for chapter_key, chapter_entry in raw_chapter_meta.items():
if isinstance(chapter_entry, Mapping):
normalized[str(chapter_key)] = normalize_chapter_meta_entry(chapter_entry)
return normalized
def get_chapter_meta_entry(state: Mapping[str, Any], chapter: int) -> Dict[str, Any]:
chapter_meta = state.get("chapter_meta", {})
if not isinstance(chapter_meta, Mapping):
return {}
for lookup_key in (f"{chapter:04d}", str(chapter)):
value = chapter_meta.get(lookup_key)
if isinstance(value, Mapping):
return normalize_chapter_meta_entry(value)
for raw_key, raw_value in chapter_meta.items():
if to_positive_int(raw_key) == chapter and isinstance(raw_value, Mapping):
return normalize_chapter_meta_entry(raw_value)
return {}
def normalize_state_runtime_sections(state: Dict[str, Any]) -> Dict[str, Any]:
if not isinstance(state, dict):
return {}
plot_threads = state.get("plot_threads")
if not isinstance(plot_threads, dict):
plot_threads = {}
state["plot_threads"] = plot_threads
plot_threads["foreshadowing"] = normalize_foreshadowing_list(plot_threads.get("foreshadowing"))
state["chapter_meta"] = normalize_chapter_meta(state.get("chapter_meta", {}))
return state
+426
View File
@@ -0,0 +1,426 @@
#!/usr/bin/env python3
# -*- coding: utf-8 -*-
"""
Style Sampler - 风格样本管理模块
管理高质量章节片段作为风格参考:
- 风格样本存储
- 按场景类型分类
- 样本选择策略
"""
import json
import sqlite3
import time
from pathlib import Path
from typing import Dict, List, Optional, Any
from dataclasses import dataclass, asdict
from datetime import datetime
from enum import Enum
from contextlib import contextmanager
from .config import get_config
from .observability import safe_append_perf_timing, safe_log_tool_call
class SceneType(Enum):
"""场景类型"""
BATTLE = "战斗"
DIALOGUE = "对话"
DESCRIPTION = "描写"
TRANSITION = "过渡"
EMOTION = "情感"
TENSION = "紧张"
COMEDY = "轻松"
@dataclass
class StyleSample:
"""风格样本"""
id: str
chapter: int
scene_type: str
content: str
score: float
tags: List[str]
created_at: str = ""
class StyleSampler:
"""风格样本管理器"""
def __init__(self, config=None):
self.config = config or get_config()
self._init_db()
def _init_db(self):
"""初始化数据库"""
self.config.ensure_dirs()
with self._get_conn() as conn:
cursor = conn.cursor()
cursor.execute("""
CREATE TABLE IF NOT EXISTS samples (
id TEXT PRIMARY KEY,
chapter INTEGER,
scene_type TEXT,
content TEXT,
score REAL,
tags TEXT,
created_at TIMESTAMP DEFAULT CURRENT_TIMESTAMP
)
""")
cursor.execute("CREATE INDEX IF NOT EXISTS idx_samples_type ON samples(scene_type)")
cursor.execute("CREATE INDEX IF NOT EXISTS idx_samples_score ON samples(score DESC)")
conn.commit()
@contextmanager
def _get_conn(self):
"""获取数据库连接(确保关闭,避免 Windows 下文件句柄泄漏导致无法清理临时目录)"""
db_path = self.config.noma_dir / "style_samples.db"
conn = sqlite3.connect(str(db_path))
try:
yield conn
finally:
conn.close()
# ==================== 样本管理 ====================
def add_sample(self, sample: StyleSample) -> bool:
"""添加风格样本"""
with self._get_conn() as conn:
cursor = conn.cursor()
try:
cursor.execute("""
INSERT INTO samples
(id, chapter, scene_type, content, score, tags, created_at)
VALUES (?, ?, ?, ?, ?, ?, ?)
""", (
sample.id,
sample.chapter,
sample.scene_type,
sample.content,
sample.score,
json.dumps(sample.tags, ensure_ascii=False),
sample.created_at or datetime.now().isoformat()
))
conn.commit()
return True
except sqlite3.IntegrityError:
return False
def get_samples_by_type(
self,
scene_type: str,
limit: int = 5,
min_score: float = 0.0
) -> List[StyleSample]:
"""按场景类型获取样本"""
with self._get_conn() as conn:
cursor = conn.cursor()
cursor.execute("""
SELECT id, chapter, scene_type, content, score, tags, created_at
FROM samples
WHERE scene_type = ? AND score >= ?
ORDER BY score DESC
LIMIT ?
""", (scene_type, min_score, limit))
return [self._row_to_sample(row) for row in cursor.fetchall()]
def get_best_samples(self, limit: int = 10) -> List[StyleSample]:
"""获取最高分样本"""
with self._get_conn() as conn:
cursor = conn.cursor()
cursor.execute("""
SELECT id, chapter, scene_type, content, score, tags, created_at
FROM samples
ORDER BY score DESC
LIMIT ?
""", (limit,))
return [self._row_to_sample(row) for row in cursor.fetchall()]
def _row_to_sample(self, row) -> StyleSample:
"""将数据库行转换为样本对象"""
return StyleSample(
id=row[0],
chapter=row[1],
scene_type=row[2],
content=row[3],
score=row[4],
tags=json.loads(row[5]) if row[5] else [],
created_at=row[6]
)
# ==================== 样本提取 ====================
def extract_candidates(
self,
chapter: int,
content: str,
review_score: float,
scenes: List[Dict]
) -> List[StyleSample]:
"""
从章节中提取风格样本候选
只有高分章节 (review_score >= 80) 才提取样本
"""
if review_score < 80:
return []
candidates = []
for scene in scenes:
scene_type = self._classify_scene_type(scene)
scene_content = scene.get("content", "")
# 跳过过短的场景
if len(scene_content) < 200:
continue
# 创建样本
sample = StyleSample(
id=f"ch{chapter}_s{scene.get('index', 0)}",
chapter=chapter,
scene_type=scene_type,
content=scene_content[:2000], # 限制长度
score=review_score / 100.0,
tags=self._extract_tags(scene_content)
)
candidates.append(sample)
return candidates
def _classify_scene_type(self, scene: Dict) -> str:
"""分类场景类型"""
summary = scene.get("summary", "").lower()
content = scene.get("content", "").lower()
# 简单关键词分类
battle_keywords = ["战斗", "攻击", "出手", "拳", "剑", "杀", "打", "斗"]
dialogue_keywords = ["说道", "问道", "笑道", "冷声", "对话"]
emotion_keywords = ["心中", "感觉", "情", "泪", "痛", "喜"]
tension_keywords = ["危险", "紧张", "恐惧", "压力"]
text = summary + content
if any(kw in text for kw in battle_keywords):
return SceneType.BATTLE.value
elif any(kw in text for kw in tension_keywords):
return SceneType.TENSION.value
elif any(kw in text for kw in dialogue_keywords):
return SceneType.DIALOGUE.value
elif any(kw in text for kw in emotion_keywords):
return SceneType.EMOTION.value
else:
return SceneType.DESCRIPTION.value
def _extract_tags(self, content: str) -> List[str]:
"""提取内容标签"""
tags = []
# 简单标签提取
if "战斗" in content or "攻击" in content:
tags.append("战斗")
if "修炼" in content or "突破" in content:
tags.append("修炼")
if "对话" in content or "说道" in content:
tags.append("对话")
if "描写" in content or "景色" in content:
tags.append("描写")
return tags[:5]
# ==================== 样本选择 ====================
def select_samples_for_chapter(
self,
chapter_outline: str,
target_types: List[str] = None,
max_samples: int = 3
) -> List[StyleSample]:
"""
为章节写作选择合适的风格样本
基于大纲分析需要什么类型的样本
"""
if target_types is None:
# 根据大纲推断需要的场景类型
target_types = self._infer_scene_types(chapter_outline)
samples = []
per_type = max(1, max_samples // len(target_types)) if target_types else max_samples
for scene_type in target_types:
type_samples = self.get_samples_by_type(scene_type, limit=per_type, min_score=0.8)
samples.extend(type_samples)
return samples[:max_samples]
def _infer_scene_types(self, outline: str) -> List[str]:
"""从大纲推断需要的场景类型"""
types = []
if any(kw in outline for kw in ["战斗", "对决", "比试", "交手"]):
types.append(SceneType.BATTLE.value)
if any(kw in outline for kw in ["对话", "谈话", "商议", "讨论"]):
types.append(SceneType.DIALOGUE.value)
if any(kw in outline for kw in ["情感", "感情", "心理"]):
types.append(SceneType.EMOTION.value)
if not types:
types = [SceneType.DESCRIPTION.value]
return types
# ==================== 统计 ====================
def get_stats(self) -> Dict[str, Any]:
"""获取样本统计"""
with self._get_conn() as conn:
cursor = conn.cursor()
cursor.execute("SELECT COUNT(*) FROM samples")
total = cursor.fetchone()[0]
cursor.execute("""
SELECT scene_type, COUNT(*) as count
FROM samples
GROUP BY scene_type
""")
by_type = {row[0]: row[1] for row in cursor.fetchall()}
cursor.execute("SELECT AVG(score) FROM samples")
avg_score = cursor.fetchone()[0] or 0
return {
"total": total,
"by_type": by_type,
"avg_score": round(avg_score, 3)
}
# ==================== CLI 接口 ====================
def main():
import argparse
import sys
from .cli_output import print_success, print_error
from .cli_args import normalize_global_project_root, load_json_arg
from .index_manager import IndexManager
parser = argparse.ArgumentParser(description="Style Sampler CLI")
parser.add_argument("--project-root", type=str, help="项目根目录")
subparsers = parser.add_subparsers(dest="command")
# 获取统计
subparsers.add_parser("stats")
# 列出样本
list_parser = subparsers.add_parser("list")
list_parser.add_argument("--type", help="按类型过滤")
list_parser.add_argument("--limit", type=int, default=10)
# 提取样本
extract_parser = subparsers.add_parser("extract")
extract_parser.add_argument("--chapter", type=int, required=True)
extract_parser.add_argument("--score", type=float, required=True)
extract_parser.add_argument("--scenes", required=True, help="JSON 格式的场景列表")
# 选择样本
select_parser = subparsers.add_parser("select")
select_parser.add_argument("--outline", required=True, help="章节大纲")
select_parser.add_argument("--max", type=int, default=3)
argv = normalize_global_project_root(sys.argv[1:])
args = parser.parse_args(argv)
command_started_at = time.perf_counter()
# 初始化
config = None
if args.project_root:
# 允许传入“工作区根目录”,统一解析到真正的 book project_root(必须包含 .noma/state.json)
from project_locator import resolve_project_root
from .config import DataModulesConfig
resolved_root = resolve_project_root(args.project_root)
config = DataModulesConfig.from_project_root(resolved_root)
sampler = StyleSampler(config)
logger = IndexManager(config)
tool_name = f"style_sampler:{args.command or 'unknown'}"
def _append_timing(success: bool, *, error_code: str | None = None, error_message: str | None = None, chapter: int | None = None):
elapsed_ms = int((time.perf_counter() - command_started_at) * 1000)
safe_append_perf_timing(
sampler.config.project_root,
tool_name=tool_name,
success=success,
elapsed_ms=elapsed_ms,
chapter=chapter,
error_code=error_code,
error_message=error_message,
)
def emit_success(data=None, message: str = "ok", chapter: int | None = None):
print_success(data, message=message)
safe_log_tool_call(logger, tool_name=tool_name, success=True)
_append_timing(True, chapter=chapter)
def emit_error(code: str, message: str, suggestion: str | None = None, chapter: int | None = None):
print_error(code, message, suggestion=suggestion)
safe_log_tool_call(
logger,
tool_name=tool_name,
success=False,
error_code=code,
error_message=message,
)
_append_timing(False, error_code=code, error_message=message, chapter=chapter)
if args.command == "stats":
stats = sampler.get_stats()
emit_success(stats, message="stats")
elif args.command == "list":
if args.type:
samples = sampler.get_samples_by_type(args.type, args.limit)
else:
samples = sampler.get_best_samples(args.limit)
emit_success([s.__dict__ for s in samples], message="samples")
elif args.command == "extract":
scenes = load_json_arg(args.scenes)
candidates = sampler.extract_candidates(
chapter=args.chapter,
content="",
review_score=args.score,
scenes=scenes,
)
added = []
skipped = []
for c in candidates:
if sampler.add_sample(c):
added.append(c.id)
else:
skipped.append(c.id)
emit_success({"added": added, "skipped": skipped}, message="extracted", chapter=args.chapter)
elif args.command == "select":
samples = sampler.select_samples_for_chapter(args.outline, max_samples=args.max)
emit_success([s.__dict__ for s in samples], message="selected")
else:
emit_error("UNKNOWN_COMMAND", "未指定有效命令", suggestion="请查看 --help")
if __name__ == "__main__":
main()
@@ -0,0 +1 @@
# data_modules tests package
@@ -0,0 +1,485 @@
#!/usr/bin/env python3
# -*- coding: utf-8 -*-
"""
API Client tests
"""
import asyncio
import json
import pytest
from data_modules.config import DataModulesConfig
from data_modules.api_client import (
EmbeddingAPIClient,
RerankAPIClient,
ModalAPIClient,
get_client,
)
class FakeResponse:
def __init__(self, status, json_data=None, text_data=""):
self.status = status
self._json = json_data
if text_data:
self._text = text_data
elif json_data is not None:
self._text = json.dumps(json_data, ensure_ascii=False)
else:
self._text = ""
async def __aenter__(self):
return self
async def __aexit__(self, exc_type, exc, tb):
return False
async def json(self):
return self._json
async def text(self):
return self._text
class FakeSession:
def __init__(self, responses):
self._responses = list(responses)
self.closed = False
def post(self, *args, **kwargs):
if not self._responses:
raise AssertionError("No more responses")
resp = self._responses.pop(0)
if isinstance(resp, Exception):
raise resp
return resp
async def close(self):
self.closed = True
@pytest.mark.asyncio
async def test_embedding_client_success_and_retry(tmp_path, monkeypatch):
config = DataModulesConfig.from_project_root(tmp_path)
config.embed_api_type = "openai"
config.api_max_retries = 2
client = EmbeddingAPIClient(config)
responses = [
FakeResponse(500, text_data="err"),
FakeResponse(
200,
json_data={
"data": [
{"embedding": [0.1, 0.2], "index": 1},
{"embedding": [0.3, 0.4], "index": 0},
]
},
),
]
fake_session = FakeSession(responses)
async def fake_get_session():
return fake_session
monkeypatch.setattr(client, "_get_session", fake_get_session)
result = await client.embed(["a", "b"])
assert result == [[0.3, 0.4], [0.1, 0.2]]
assert client.stats.total_calls == 1
assert client.stats.errors == 0
@pytest.mark.asyncio
async def test_embedding_client_timeout_and_error(tmp_path, monkeypatch):
config = DataModulesConfig.from_project_root(tmp_path)
config.embed_api_type = "openai"
config.api_max_retries = 1
client = EmbeddingAPIClient(config)
responses = [asyncio.TimeoutError()]
fake_session = FakeSession(responses)
async def fake_get_session():
return fake_session
monkeypatch.setattr(client, "_get_session", fake_get_session)
result = await client.embed(["x"])
assert result is None
assert client.stats.errors == 1
@pytest.mark.asyncio
async def test_embedding_batch(tmp_path, monkeypatch):
config = DataModulesConfig.from_project_root(tmp_path)
config.embed_batch_size = 2
client = EmbeddingAPIClient(config)
async def fake_embed(texts):
if len(texts) == 2:
return [[1.0, 0.0], [0.0, 1.0]]
return None
monkeypatch.setattr(client, "embed", fake_embed)
result = await client.embed_batch(["a", "b", "c"], skip_failures=True)
assert result[0] is not None
assert result[2] is None
result_fail = await client.embed_batch(["a", "b", "c"], skip_failures=False)
assert result_fail == []
def test_embedding_build_url_and_payload(tmp_path):
config = DataModulesConfig.from_project_root(tmp_path)
config.embed_api_type = "openai"
config.embed_base_url = "https://api.example.com"
client = EmbeddingAPIClient(config)
assert client._build_url().endswith("/v1/embeddings")
payload = client._build_payload(["hi"])
assert payload["model"] == config.embed_model
config.embed_base_url = "https://api.example.com/v1"
assert client._build_url().endswith("/v1/embeddings")
config.embed_base_url = "https://api.example.com/v1/embeddings"
assert client._build_url().endswith("/v1/embeddings")
config.embed_api_type = "modal"
config.embed_base_url = "https://modal.example.com/embed"
assert client._build_url() == "https://modal.example.com/embed"
payload = client._build_payload(["hi"])
assert "encoding_format" not in payload
@pytest.mark.asyncio
async def test_rerank_client_success(tmp_path, monkeypatch):
config = DataModulesConfig.from_project_root(tmp_path)
config.rerank_api_type = "openai"
config.api_max_retries = 1
client = RerankAPIClient(config)
responses = [
FakeResponse(
200,
json_data={"results": [{"index": 0, "relevance_score": 0.9}]},
)
]
fake_session = FakeSession(responses)
async def fake_get_session():
return fake_session
monkeypatch.setattr(client, "_get_session", fake_get_session)
result = await client.rerank("q", ["doc1"], top_n=1)
assert result[0]["index"] == 0
assert client.stats.total_calls == 1
@pytest.mark.asyncio
async def test_rerank_retry_and_empty(tmp_path, monkeypatch):
config = DataModulesConfig.from_project_root(tmp_path)
config.rerank_api_type = "openai"
config.api_max_retries = 2
client = RerankAPIClient(config)
responses = [
FakeResponse(503, text_data="err"),
FakeResponse(
200,
json_data={"results": [{"index": 0, "relevance_score": 0.8}]},
),
]
fake_session = FakeSession(responses)
async def fake_get_session():
return fake_session
monkeypatch.setattr(client, "_get_session", fake_get_session)
result = await client.rerank("q", ["doc1"], top_n=1)
assert result[0]["relevance_score"] == 0.8
assert await client.rerank("q", []) == []
@pytest.mark.asyncio
async def test_modal_client_warmup_and_passthrough(tmp_path, monkeypatch):
config = DataModulesConfig.from_project_root(tmp_path)
client = ModalAPIClient(config)
async def fake_warmup():
return None
async def fake_embed(texts):
return [[0.1, 0.2] for _ in texts]
async def fake_rerank(query, documents, top_n=None):
return [{"index": 0, "relevance_score": 1.0}]
monkeypatch.setattr(client._embed_client, "warmup", fake_warmup)
monkeypatch.setattr(client._rerank_client, "warmup", fake_warmup)
monkeypatch.setattr(client._embed_client, "embed", fake_embed)
monkeypatch.setattr(client._rerank_client, "rerank", fake_rerank)
await client.warmup()
assert client._warmed_up["embed"] is True
assert client._warmed_up["rerank"] is True
emb = await client.embed(["hi"])
assert emb[0] == [0.1, 0.2]
rr = await client.rerank("q", ["doc"])
assert rr[0]["index"] == 0
def test_get_client_singleton(tmp_path):
cfg = DataModulesConfig.from_project_root(tmp_path)
client1 = get_client(cfg)
client2 = get_client()
assert client1 is client2
client3 = get_client(cfg)
assert client3 is not client1
@pytest.mark.asyncio
async def test_embedding_empty_and_error_paths(tmp_path, monkeypatch):
config = DataModulesConfig.from_project_root(tmp_path)
config.embed_api_key = "sk-test"
config.api_max_retries = 1
client = EmbeddingAPIClient(config)
assert await client.embed([]) == []
headers = client._build_headers()
assert headers["Authorization"] == "Bearer sk-test"
fake_session = FakeSession([FakeResponse(400, text_data="bad request")])
async def fake_get_session():
return fake_session
monkeypatch.setattr(client, "_get_session", fake_get_session)
result = await client.embed(["x"])
assert result is None
assert client.stats.errors == 1
@pytest.mark.asyncio
async def test_embedding_exception_and_close(tmp_path, monkeypatch):
config = DataModulesConfig.from_project_root(tmp_path)
config.api_max_retries = 1
client = EmbeddingAPIClient(config)
class BoomSession:
def __init__(self):
self.closed = False
def post(self, *args, **kwargs):
raise RuntimeError("boom")
async def close(self):
self.closed = True
session = BoomSession()
async def fake_get_session():
return session
monkeypatch.setattr(client, "_get_session", fake_get_session)
result = await client.embed(["x"])
assert result is None
assert client.stats.errors == 1
client._session = session
await client.close()
assert session.closed is True
def test_rerank_headers_payload_and_stats(tmp_path, capsys):
config = DataModulesConfig.from_project_root(tmp_path)
config.rerank_api_key = "rk-test"
client = RerankAPIClient(config)
headers = client._build_headers()
assert headers["Authorization"] == "Bearer rk-test"
payload = client._build_payload("q", ["doc"], top_n=2)
assert payload["top_n"] == 2
modal = ModalAPIClient(config)
modal._embed_client.stats.total_calls = 1
modal._embed_client.stats.total_time = 2.0
modal.print_stats()
output = capsys.readouterr().out
assert "EMBED" in output
@pytest.mark.asyncio
async def test_rerank_non_retry_error(tmp_path, monkeypatch):
config = DataModulesConfig.from_project_root(tmp_path)
config.api_max_retries = 1
client = RerankAPIClient(config)
fake_session = FakeSession([FakeResponse(400, text_data="bad request")])
async def fake_get_session():
return fake_session
monkeypatch.setattr(client, "_get_session", fake_get_session)
result = await client.rerank("q", ["doc"])
assert result is None
assert client.stats.errors == 1
@pytest.mark.asyncio
async def test_embedding_session_parse_and_retry_paths(tmp_path, monkeypatch):
config = DataModulesConfig.from_project_root(tmp_path)
config.embed_api_type = "modal"
config.api_max_retries = 2
config.api_retry_delay = 0
client = EmbeddingAPIClient(config)
session = await client._get_session()
assert session is not None
await client.close()
assert client._parse_response({}) is None
parsed = client._parse_response({"data": [{"embedding": [1.0, 2.0]}]})
assert parsed == [[1.0, 2.0]]
responses = [
asyncio.TimeoutError(),
FakeResponse(200, text_data=json.dumps({"data": [{"embedding": [0.1], "index": 0}]})),
]
fake_session = FakeSession(responses)
async def fake_get_session():
return fake_session
monkeypatch.setattr(client, "_get_session", fake_get_session)
result = await client.embed(["x"])
assert result == [[0.1]]
@pytest.mark.asyncio
async def test_embedding_exception_retry_and_batch(tmp_path, monkeypatch):
config = DataModulesConfig.from_project_root(tmp_path)
config.api_max_retries = 2
config.api_retry_delay = 0
client = EmbeddingAPIClient(config)
responses = [
RuntimeError("boom"),
FakeResponse(200, text_data=json.dumps({"data": [{"embedding": [0.2], "index": 0}]})),
]
fake_session = FakeSession(responses)
async def fake_get_session():
return fake_session
monkeypatch.setattr(client, "_get_session", fake_get_session)
result = await client.embed(["x"])
assert result == [[0.2]]
assert await client.embed_batch([]) == []
async def fake_embed(texts):
return [[0.0] for _ in texts]
monkeypatch.setattr(client, "embed", fake_embed)
await client.warmup()
assert client._warmed_up is True
@pytest.mark.asyncio
async def test_rerank_modal_retry_and_warmup(tmp_path, monkeypatch):
config = DataModulesConfig.from_project_root(tmp_path)
config.rerank_api_type = "modal"
config.rerank_base_url = "https://modal.example.com/rerank"
config.api_max_retries = 2
config.api_retry_delay = 0
client = RerankAPIClient(config)
session = await client._get_session()
assert session is not None
await client.close()
payload = client._build_payload("q", ["doc"], top_n=1)
assert payload["top_n"] == 1
assert client._build_url() == "https://modal.example.com/rerank"
assert client._parse_response({"results": [{"index": 0}]}) == [{"index": 0}]
responses = [
asyncio.TimeoutError(),
FakeResponse(200, json_data={"results": [{"index": 0, "relevance_score": 1.0}]}),
]
fake_session = FakeSession(responses)
async def fake_get_session():
return fake_session
monkeypatch.setattr(client, "_get_session", fake_get_session)
result = await client.rerank("q", ["doc"])
assert result[0]["index"] == 0
responses = [
RuntimeError("boom"),
FakeResponse(200, json_data={"results": [{"index": 0, "relevance_score": 0.5}]}),
]
fake_session = FakeSession(responses)
async def fake_get_session2():
return fake_session
monkeypatch.setattr(client, "_get_session", fake_get_session2)
result = await client.rerank("q", ["doc"])
assert result[0]["relevance_score"] == 0.5
async def fake_rerank(query, docs, top_n=None):
return [{"index": 0, "relevance_score": 1.0}]
monkeypatch.setattr(client, "rerank", fake_rerank)
await client.warmup()
assert client._warmed_up is True
@pytest.mark.asyncio
async def test_modal_client_helpers(tmp_path, monkeypatch, capsys):
config = DataModulesConfig.from_project_root(tmp_path)
client = ModalAPIClient(config)
async def fake_embed_batch(texts, skip_failures=True):
return [[0.1] for _ in texts]
monkeypatch.setattr(client._embed_client, "embed_batch", fake_embed_batch)
result = await client.embed_batch(["a", "b"])
assert result[0] == [0.1]
async def fail_warmup():
raise RuntimeError("fail")
async def ok_warmup():
return None
monkeypatch.setattr(client, "_warmup_embed", fail_warmup)
monkeypatch.setattr(client, "_warmup_rerank", ok_warmup)
await client.warmup()
output = capsys.readouterr().out
assert "[FAIL]" in output
async def fake_get_session():
return FakeSession([])
monkeypatch.setattr(client._embed_client, "_get_session", fake_get_session)
session = await client._get_session()
assert session is not None
closed = {"embed": False, "rerank": False}
async def close_embed():
closed["embed"] = True
async def close_rerank():
closed["rerank"] = True
monkeypatch.setattr(client._embed_client, "close", close_embed)
monkeypatch.setattr(client._rerank_client, "close", close_rerank)
await client.close()
assert closed["embed"] and closed["rerank"]
@@ -0,0 +1,74 @@
#!/usr/bin/env python3
# -*- coding: utf-8 -*-
from pathlib import Path
import pytest
def _load_archive_module():
import sys
scripts_dir = Path(__file__).resolve().parents[2]
if str(scripts_dir) not in sys.path:
sys.path.insert(0, str(scripts_dir))
import archive_manager
return archive_manager
@pytest.fixture
def archive_env(tmp_path):
noma = tmp_path / ".noma"
noma.mkdir(parents=True, exist_ok=True)
state_path = noma / "state.json"
state_path.write_text(
'{"progress":{"current_chapter":10},"plot_threads":{},"review_checkpoints":[]}',
encoding="utf-8",
)
return tmp_path
def test_archive_remove_from_state_missing_sections(archive_env):
module = _load_archive_module()
manager = module.ArchiveManager(project_root=archive_env)
state = {
"progress": {"current_chapter": 50},
}
updated = manager.remove_from_state(state, inactive_chars=[], resolved_threads=[], old_reviews=[])
assert updated.get("progress", {}).get("current_chapter") == 50
def test_archive_check_trigger_conditions_edges(archive_env):
module = _load_archive_module()
manager = module.ArchiveManager(project_root=archive_env)
manager.config["chapter_trigger"] = 10
manager.config["file_size_trigger_mb"] = 9999.0
trigger = manager.check_trigger_conditions({"progress": {"current_chapter": 20}})
assert trigger["chapter_trigger"] is True
assert trigger["should_archive"] is True
def test_archive_identify_old_reviews_handles_mixed_formats(archive_env):
module = _load_archive_module()
manager = module.ArchiveManager(project_root=archive_env)
manager.config["review_old_threshold"] = 5
state = {
"progress": {"current_chapter": 30},
"review_checkpoints": [
{"chapters": "20-22", "report": "r1.md"},
{"chapter_range": [10, 12], "date": "2026-01-01"},
{"report": "Review_Ch5-6.md"},
],
}
results = manager.identify_old_reviews(state)
assert len(results) == 3
assert all(row["chapters_since_review"] >= 5 for row in results)
@@ -0,0 +1,51 @@
#!/usr/bin/env python3
# -*- coding: utf-8 -*-
import sys
from pathlib import Path
def _load_module():
scripts_dir = Path(__file__).resolve().parents[2]
if str(scripts_dir) not in sys.path:
sys.path.insert(0, str(scripts_dir))
import chapter_paths
return chapter_paths
def test_default_chapter_draft_path_uses_outline_heading_title(tmp_path):
module = _load_module()
outline_dir = tmp_path / "大纲"
outline_dir.mkdir(parents=True, exist_ok=True)
(outline_dir / "第1卷-详细大纲.md").write_text("### 第1章:测试标题\n测试大纲", encoding="utf-8")
draft_path = module.default_chapter_draft_path(tmp_path, 1)
assert draft_path.name == "第0001章-测试标题.md"
def test_default_chapter_draft_path_falls_back_to_split_outline_filename(tmp_path):
module = _load_module()
outline_dir = tmp_path / "大纲"
outline_dir.mkdir(parents=True, exist_ok=True)
(outline_dir / "第0002章-标题 文件.md").write_text("无章节标题 heading", encoding="utf-8")
draft_path = module.default_chapter_draft_path(tmp_path, 2)
assert draft_path.name == "第0002章-标题_文件.md"
def test_find_chapter_file_supports_titled_flat_filename(tmp_path):
module = _load_module()
chapter_path = tmp_path / "正文" / "第0003章-山雨欲来.md"
chapter_path.parent.mkdir(parents=True, exist_ok=True)
chapter_path.write_text("正文", encoding="utf-8")
found = module.find_chapter_file(tmp_path, 3)
assert found == chapter_path
@@ -0,0 +1,62 @@
#!/usr/bin/env python3
# -*- coding: utf-8 -*-
"""
Config tests
"""
import os
from data_modules import config as config_module
from data_modules.config import DataModulesConfig, get_config, set_project_root
def test_config_paths_and_defaults(tmp_path):
cfg = DataModulesConfig.from_project_root(tmp_path)
assert cfg.project_root == tmp_path
assert cfg.noma_dir.name == ".noma"
assert cfg.state_file.name == "state.json"
assert cfg.index_db.name == "index.db"
assert cfg.rag_db.name == "rag.db"
assert cfg.vector_db.name == "vectors.db"
cfg.ensure_dirs()
assert cfg.noma_dir.exists()
def test_get_config_and_set_project_root(tmp_path):
set_project_root(tmp_path)
cfg = get_config()
assert cfg.project_root == tmp_path
def test_load_dotenv(monkeypatch, tmp_path):
# prepare .env
env_path = tmp_path / ".env"
env_path.write_text("EMBED_BASE_URL=https://example.com\n", encoding="utf-8")
monkeypatch.chdir(tmp_path)
monkeypatch.delenv("EMBED_BASE_URL", raising=False)
# call loader explicitly
config_module._load_dotenv()
assert os.environ.get("EMBED_BASE_URL") == "https://example.com"
def test_config_default_context_template_weights_dynamic_is_available(tmp_path):
cfg = DataModulesConfig.from_project_root(tmp_path)
dynamic = cfg.context_template_weights_dynamic
assert isinstance(dynamic, dict)
assert "early" in dynamic
assert "mid" in dynamic
assert "late" in dynamic
assert "plot" in dynamic["early"]
def test_config_dynamic_template_weights_are_independent_instances(tmp_path):
cfg1 = DataModulesConfig.from_project_root(tmp_path)
cfg2 = DataModulesConfig.from_project_root(tmp_path)
cfg1.context_template_weights_dynamic["early"]["plot"]["core"] = 0.77
assert cfg2.context_template_weights_dynamic["early"]["plot"]["core"] != 0.77
@@ -0,0 +1,657 @@
#!/usr/bin/env python3
# -*- coding: utf-8 -*-
"""
ContextManager and SnapshotManager tests
"""
import json
import logging
import pytest
from data_modules.config import DataModulesConfig
from data_modules.index_manager import (
IndexManager,
EntityMeta,
ChapterReadingPowerMeta,
ReviewMetrics,
)
from data_modules.context_manager import ContextManager
from data_modules.snapshot_manager import SnapshotManager, SnapshotVersionMismatch
from data_modules.query_router import QueryRouter
@pytest.fixture
def temp_project(tmp_path):
cfg = DataModulesConfig.from_project_root(tmp_path)
cfg.ensure_dirs()
return cfg
def test_snapshot_manager_roundtrip(temp_project):
manager = SnapshotManager(temp_project)
payload = {"hello": "world"}
manager.save_snapshot(1, payload)
loaded = manager.load_snapshot(1)
assert loaded["payload"] == payload
def test_snapshot_version_mismatch(temp_project):
manager = SnapshotManager(temp_project, version="1.0")
manager.save_snapshot(1, {"a": 1})
other = SnapshotManager(temp_project, version="2.0")
with pytest.raises(SnapshotVersionMismatch):
other.load_snapshot(1)
def test_snapshot_delete_roundtrip(temp_project):
manager = SnapshotManager(temp_project)
manager.save_snapshot(2, {"x": 1})
assert manager.delete_snapshot(2) is True
assert manager.load_snapshot(2) is None
def test_context_manager_build_and_filter(temp_project):
state = {
"protagonist_state": {"name": "萧炎", "location": {"current": "天云宗"}},
"chapter_meta": {"0001": {"hook": "测试"}},
}
temp_project.state_file.write_text(json.dumps(state, ensure_ascii=False), encoding="utf-8")
# preferences and memory
(temp_project.noma_dir / "preferences.json").write_text(json.dumps({"tone": "热血"}, ensure_ascii=False), encoding="utf-8")
(temp_project.noma_dir / "project_memory.json").write_text(json.dumps({"patterns": []}, ensure_ascii=False), encoding="utf-8")
idx = IndexManager(temp_project)
idx.upsert_entity(
EntityMeta(
id="xiaoyan",
type="角色",
canonical_name="萧炎",
current={},
first_appearance=1,
last_appearance=1,
)
)
idx.upsert_entity(
EntityMeta(
id="bad",
type="角色",
canonical_name="坏人",
current={},
first_appearance=1,
last_appearance=1,
)
)
idx.record_appearance("xiaoyan", 1, ["萧炎"], 1.0)
idx.record_appearance("bad", 1, ["坏人"], 1.0)
invalid_id = idx.mark_invalid_fact("entity", "bad", "错误")
idx.resolve_invalid_fact(invalid_id, "confirm")
manager = ContextManager(temp_project)
payload = manager.build_context(1, use_snapshot=False, save_snapshot=False)
characters = payload["sections"]["scene"]["content"]["appearing_characters"]
assert any(c.get("entity_id") == "xiaoyan" for c in characters)
assert not any(c.get("entity_id") == "bad" for c in characters)
assert payload["sections"]["preferences"]["content"].get("tone") == "热血"
def test_context_manager_loads_volume_outline_file(temp_project):
state = {
"progress": {
"volumes_planned": [
{"volume": 1, "chapters_range": "1-10"},
]
},
"protagonist_state": {},
"chapter_meta": {},
"disambiguation_warnings": [],
"disambiguation_pending": [],
}
temp_project.state_file.write_text(json.dumps(state, ensure_ascii=False), encoding="utf-8")
temp_project.outline_dir.mkdir(parents=True, exist_ok=True)
(temp_project.outline_dir / "第1卷-详细大纲.md").write_text(
"### 第2章:测试标题\n测试大纲\n\n### 第3章:下一章",
encoding="utf-8",
)
manager = ContextManager(temp_project)
payload = manager.build_context(2, use_snapshot=False, save_snapshot=False)
outline = payload["sections"]["core"]["content"]["chapter_outline"]
assert "### 第2章:测试标题" in outline
assert "测试大纲" in outline
def test_query_router():
router = QueryRouter()
assert router.route("角色是谁") == "entity"
assert router.route("发生了什么剧情") == "plot"
intent = router.route_intent("第10-20章萧炎和药老关系图谱")
assert intent["intent"] == "relationship"
assert intent["needs_graph"] is True
assert intent["time_scope"]["from_chapter"] == 10
assert intent["time_scope"]["to_chapter"] == 20
plans = router.plan_subqueries(intent)
assert plans
assert plans[0]["strategy"] in {"graph_lookup", "graph_hybrid"}
assert "A" in router.split("A, B;C")
def test_context_snapshot_respects_template(temp_project):
state = {
"protagonist_state": {"name": "萧炎"},
"chapter_meta": {},
"disambiguation_warnings": [],
"disambiguation_pending": [],
}
temp_project.state_file.write_text(json.dumps(state, ensure_ascii=False), encoding="utf-8")
manager = ContextManager(temp_project)
plot_payload = manager.build_context(1, template="plot", use_snapshot=True, save_snapshot=True)
battle_payload = manager.build_context(1, template="battle", use_snapshot=True, save_snapshot=True)
assert plot_payload.get("template") == "plot"
assert battle_payload.get("template") == "battle"
def test_context_manager_applies_ranker_and_contract_meta(temp_project):
state = {
"protagonist_state": {"name": "萧炎"},
"chapter_meta": {
"0002": {"hook": "平稳"},
"0003": {"hook": "留下悬念"},
},
"disambiguation_warnings": [
{"chapter": 1, "message": "普通告警"},
{"chapter": 3, "message": "critical 冲突告警", "severity": "high"},
],
"disambiguation_pending": [],
}
temp_project.state_file.write_text(json.dumps(state, ensure_ascii=False), encoding="utf-8")
manager = ContextManager(temp_project)
payload = manager.build_context(4, use_snapshot=False, save_snapshot=False)
assert payload["meta"].get("context_contract_version") == "v2"
recent_meta = payload["sections"]["core"]["content"]["recent_meta"]
if recent_meta:
assert recent_meta[0]["chapter"] == 3
warnings = payload["sections"]["alerts"]["content"]["disambiguation_warnings"]
if warnings and isinstance(warnings[0], dict):
assert "critical" in str(warnings[0].get("message", "")) or warnings[0].get("severity") == "high"
def test_context_manager_includes_reader_signal_and_genre_profile(temp_project):
state = {
"project": {"genre": "xuanhuan"},
"protagonist_state": {"name": "萧炎"},
"chapter_meta": {},
"disambiguation_warnings": [],
"disambiguation_pending": [],
}
temp_project.state_file.write_text(json.dumps(state, ensure_ascii=False), encoding="utf-8")
idx = IndexManager(temp_project)
idx.save_chapter_reading_power(
ChapterReadingPowerMeta(
chapter=3,
hook_type="悬念钩",
hook_strength="strong",
coolpoint_patterns=["身份掉马"],
)
)
idx.save_review_metrics(
ReviewMetrics(
start_chapter=1,
end_chapter=3,
overall_score=72,
dimension_scores={"plot": 72},
severity_counts={"high": 1},
critical_issues=["节奏拖沓"],
)
)
manager = ContextManager(temp_project)
payload = manager.build_context(4, use_snapshot=False, save_snapshot=False)
reader_signal = payload["sections"]["reader_signal"]["content"]
assert "recent_reading_power" in reader_signal
assert "pattern_usage" in reader_signal
assert "hook_type_usage" in reader_signal
assert "review_trend" in reader_signal
assert isinstance(reader_signal.get("low_score_ranges"), list)
genre_profile = payload["sections"]["genre_profile"]["content"]
assert genre_profile.get("genre") == "xuanhuan"
assert "profile_excerpt" in genre_profile
assert "taxonomy_excerpt" in genre_profile
def test_context_manager_genre_section_and_refs_extraction(temp_project):
refs_dir = temp_project.project_root / ".claude" / "references"
refs_dir.mkdir(parents=True, exist_ok=True)
(refs_dir / "genre-profiles.md").write_text(
"""
## shuangwen
- 节奏快
- 打脸密集
## xuanhuan
- 升级线清晰
- 资源争夺
""".strip(),
encoding="utf-8",
)
(refs_dir / "reading-power-taxonomy.md").write_text(
"""
## xuanhuan
- 钩子强度优先 strong
- 爽点使用战力跨级
""".strip(),
encoding="utf-8",
)
manager = ContextManager(temp_project)
profile = manager._load_genre_profile({"project": {"genre": "xuanhuan"}})
assert profile["genre"] == "xuanhuan"
assert "升级线清晰" in profile["profile_excerpt"]
assert "钩子强度" in profile["taxonomy_excerpt"]
assert isinstance(profile["reference_hints"], list)
assert profile["reference_hints"]
fallback_excerpt = manager._extract_genre_section("## a\n1\n## b\n2", "unknown")
assert fallback_excerpt.startswith("## a")
def test_context_manager_reader_signal_with_debt_and_disable_switch(temp_project):
manager = ContextManager(temp_project)
manager.config.context_reader_signal_include_debt = True
signal = manager._load_reader_signal(chapter=5)
assert "debt_summary" in signal
manager.config.context_reader_signal_enabled = False
assert manager._load_reader_signal(chapter=5) == {}
manager.config.context_genre_profile_enabled = False
assert manager._load_genre_profile({"project": {"genre": "xuanhuan"}}) == {}
def test_context_manager_includes_writing_guidance(temp_project):
state = {
"project": {"genre": "xuanhuan"},
"protagonist_state": {"name": "萧炎"},
"chapter_meta": {},
"disambiguation_warnings": [],
"disambiguation_pending": [],
}
temp_project.state_file.write_text(json.dumps(state, ensure_ascii=False), encoding="utf-8")
idx = IndexManager(temp_project)
idx.save_chapter_reading_power(
ChapterReadingPowerMeta(
chapter=3,
hook_type="悬念钩",
hook_strength="strong",
coolpoint_patterns=["身份掉马"],
)
)
idx.save_review_metrics(
ReviewMetrics(
start_chapter=1,
end_chapter=3,
overall_score=70,
dimension_scores={"plot": 70},
severity_counts={"high": 1},
critical_issues=["节奏拖沓"],
)
)
manager = ContextManager(temp_project)
payload = manager.build_context(4, use_snapshot=False, save_snapshot=False)
guidance = payload["sections"]["writing_guidance"]["content"]
assert guidance.get("chapter") == 4
items = guidance.get("guidance_items") or []
assert isinstance(items, list)
assert items
assert guidance.get("signals_used", {}).get("genre") == "xuanhuan"
checklist = guidance.get("checklist") or []
assert isinstance(checklist, list)
assert checklist
checklist_score = guidance.get("checklist_score") or {}
assert isinstance(checklist_score, dict)
assert "score" in checklist_score
assert "completion_rate" in checklist_score
first_item = checklist[0]
assert isinstance(first_item, dict)
assert {"id", "label", "weight", "required", "source", "verify_hint"}.issubset(first_item.keys())
persisted = idx.get_writing_checklist_score(4)
assert isinstance(persisted, dict)
assert persisted.get("chapter") == 4
assert persisted.get("score") is not None
def test_context_manager_dynamic_weights_and_composite_genre(temp_project):
refs_dir = temp_project.project_root / ".claude" / "references"
refs_dir.mkdir(parents=True, exist_ok=True)
(refs_dir / "genre-profiles.md").write_text(
"""
## xuanhuan
- 升级线清晰
## realistic
- 社会议题映射
""".strip(),
encoding="utf-8",
)
(refs_dir / "reading-power-taxonomy.md").write_text(
"""
## xuanhuan
- 钩子强度优先
## realistic
- 人物动机一致
""".strip(),
encoding="utf-8",
)
state = {
"project": {"genre": "xuanhuan+realistic"},
"protagonist_state": {"name": "萧炎"},
"chapter_meta": {},
"disambiguation_warnings": [],
"disambiguation_pending": [],
}
temp_project.state_file.write_text(json.dumps(state, ensure_ascii=False), encoding="utf-8")
manager = ContextManager(temp_project)
payload_early = manager.build_context(10, template="plot", use_snapshot=False, save_snapshot=False)
payload_late = manager.build_context(150, template="plot", use_snapshot=False, save_snapshot=False)
assert payload_early.get("weights", {}).get("core") >= payload_late.get("weights", {}).get("core")
assert payload_late.get("weights", {}).get("global") >= payload_early.get("weights", {}).get("global")
assert payload_early.get("meta", {}).get("context_weight_stage") == "early"
assert payload_late.get("meta", {}).get("context_weight_stage") == "late"
profile = payload_early["sections"]["genre_profile"]["content"]
assert profile.get("composite") is True
assert profile.get("genre") == "xuanhuan"
assert isinstance(profile.get("genres"), list)
assert "realistic" in (profile.get("genres") or [])
assert isinstance(profile.get("composite_hints"), list)
assert profile.get("composite_hints")
def test_context_manager_genre_alias_guidance_and_heading_extraction(temp_project):
refs_dir = temp_project.project_root / ".claude" / "references"
refs_dir.mkdir(parents=True, exist_ok=True)
(refs_dir / "genre-profiles.md").write_text(
"""
### 电竞
- 联赛升级
### 直播文
- 反馈闭环
### 克苏鲁
- 真相代价
""".strip(),
encoding="utf-8",
)
(refs_dir / "reading-power-taxonomy.md").write_text(
"""
### 电竞
- 战术决策点
""".strip(),
encoding="utf-8",
)
state = {
"project": {"genre": "电竞"},
"protagonist_state": {"name": "林燃"},
"chapter_meta": {},
"disambiguation_warnings": [],
"disambiguation_pending": [],
}
temp_project.state_file.write_text(json.dumps(state, ensure_ascii=False), encoding="utf-8")
manager = ContextManager(temp_project)
payload = manager.build_context(12, template="plot", use_snapshot=False, save_snapshot=False)
guidance = payload["sections"]["writing_guidance"]["content"]
items = guidance.get("guidance_items") or []
assert any("战术决策点" in str(text) for text in items)
assert any("网文节奏基线" in str(text) for text in items)
assert any("兑现密度基线" in str(text) for text in items)
def test_context_manager_genre_aliases_normalized_for_profile_lookup(temp_project):
refs_dir = temp_project.project_root / ".claude" / "references"
refs_dir.mkdir(parents=True, exist_ok=True)
(refs_dir / "genre-profiles.md").write_text(
"""
## 电竞
- 联赛升级
## 直播文
- 实时反馈
## 克苏鲁
- 真相代价
""".strip(),
encoding="utf-8",
)
(refs_dir / "reading-power-taxonomy.md").write_text(
"""
## 电竞
- 决策后果
## 直播文
- 数据闭环
## 克苏鲁
- 规则优先
""".strip(),
encoding="utf-8",
)
manager = ContextManager(temp_project)
assert manager._parse_genre_tokens("电竞文") == ["电竞"]
assert manager._parse_genre_tokens("直播") == ["直播文"]
assert manager._parse_genre_tokens("克系") == ["克苏鲁"]
assert manager._parse_genre_tokens("修仙/玄幻") == ["修仙"]
assert manager._parse_genre_tokens("都市修真") == ["都市异能"]
assert manager._parse_genre_tokens("古言脑洞") == ["古言"]
state = {
"project": {"genre": "电竞文+直播"},
"protagonist_state": {"name": "叶修"},
"chapter_meta": {},
"disambiguation_warnings": [],
"disambiguation_pending": [],
}
temp_project.state_file.write_text(json.dumps(state, ensure_ascii=False), encoding="utf-8")
payload = manager.build_context(20, template="plot", use_snapshot=False, save_snapshot=False)
profile = payload["sections"]["genre_profile"]["content"]
assert profile.get("genre") == "电竞"
assert "直播文" in (profile.get("genres") or [])
def test_context_manager_enables_methodology_for_xianxia(temp_project):
state = {
"project": {"genre": "修仙"},
"protagonist_state": {"name": "韩立"},
"chapter_meta": {},
"disambiguation_warnings": [],
"disambiguation_pending": [],
}
temp_project.state_file.write_text(json.dumps(state, ensure_ascii=False), encoding="utf-8")
manager = ContextManager(temp_project)
manager.config.context_writing_checklist_max_items = 8
payload = manager.build_context(21, template="plot", use_snapshot=False, save_snapshot=False)
guidance = payload["sections"]["writing_guidance"]["content"]
strategy = guidance.get("methodology") or {}
assert strategy.get("enabled") is True
assert strategy.get("pilot") == "xianxia"
assert strategy.get("genre_profile_key") == "xianxia"
assert guidance.get("signals_used", {}).get("methodology_enabled") is True
assert isinstance(strategy.get("observability"), dict)
def test_context_manager_enables_methodology_for_non_xianxia_by_default(temp_project):
state = {
"project": {"genre": "xuanhuan"},
"protagonist_state": {"name": "萧炎"},
"chapter_meta": {},
"disambiguation_warnings": [],
"disambiguation_pending": [],
}
temp_project.state_file.write_text(json.dumps(state, ensure_ascii=False), encoding="utf-8")
manager = ContextManager(temp_project)
payload = manager.build_context(21, template="plot", use_snapshot=False, save_snapshot=False)
guidance = payload["sections"]["writing_guidance"]["content"]
strategy = guidance.get("methodology") or {}
assert strategy.get("enabled") is True
assert strategy.get("genre_profile_key") == "xuanhuan"
assert guidance.get("signals_used", {}).get("methodology_enabled") is True
def test_context_manager_allows_methodology_whitelist_restriction(temp_project):
state = {
"project": {"genre": "直播文"},
"protagonist_state": {"name": "林默"},
"chapter_meta": {},
"disambiguation_warnings": [],
"disambiguation_pending": [],
}
temp_project.state_file.write_text(json.dumps(state, ensure_ascii=False), encoding="utf-8")
manager = ContextManager(temp_project)
manager.config.context_methodology_genre_whitelist = ("xianxia",)
payload = manager.build_context(21, template="plot", use_snapshot=False, save_snapshot=False)
guidance = payload["sections"]["writing_guidance"]["content"]
strategy = guidance.get("methodology") or {}
assert strategy == {}
assert guidance.get("signals_used", {}).get("methodology_enabled") is False
def test_context_manager_compact_text_truncation(temp_project):
manager = ContextManager(temp_project)
manager.config.context_compact_text_enabled = True
manager.config.context_compact_min_budget = 80
manager.config.context_compact_head_ratio = 0.6
content = {"a": "x" * 200, "b": "y" * 200}
compact = manager._compact_json_text(content, budget=120)
assert len(compact) <= 120
assert "[TRUNCATED]" in compact
manager.config.context_compact_text_enabled = False
raw_cut = manager._compact_json_text(content, budget=100)
assert len(raw_cut) <= 100
def test_context_manager_persist_writing_checklist_score_logs_failure(temp_project, monkeypatch, caplog):
manager = ContextManager(temp_project)
def _raise_save_error(_meta):
raise RuntimeError("simulated save failure")
monkeypatch.setattr(manager.index_manager, "save_writing_checklist_score", _raise_save_error)
with caplog.at_level(logging.WARNING):
manager._persist_writing_checklist_score(
{
"chapter": 6,
"score": 70.0,
"total_items": 3,
"required_items": 1,
"completed_items": 1,
"completed_required": 1,
"total_weight": 3.0,
"completed_weight": 1.0,
"completion_rate": 0.33,
"pending_items": ["test"],
}
)
message_text = "\n".join(record.getMessage() for record in caplog.records)
assert "failed to persist writing checklist score" in message_text
def test_context_manager_composite_genre_boundary_three_plus(temp_project):
manager = ContextManager(temp_project)
manager.config.context_genre_profile_support_composite = True
manager.config.context_genre_profile_max_genres = 3
genre_raw = "电竞文+直播+克系+修仙/玄幻+电竞文"
tokens = manager._parse_genre_tokens(genre_raw)
assert tokens[:4] == ["电竞", "直播文", "克苏鲁", "修仙"]
state = {
"project": {"genre": genre_raw},
"protagonist_state": {"name": "主角"},
"chapter_meta": {},
"disambiguation_warnings": [],
"disambiguation_pending": [],
}
profile = manager._load_genre_profile(state)
assert profile.get("composite") is True
assert profile.get("genres") == ["电竞", "直播文", "克苏鲁"]
assert profile.get("secondary_genres") == ["直播文", "克苏鲁"]
profile_again = manager._load_genre_profile(state)
assert profile_again.get("genres") == profile.get("genres")
def test_context_manager_dynamic_weights_from_config_override(temp_project):
manager = ContextManager(temp_project)
manager.config.context_dynamic_budget_enabled = True
manager.config.context_template_weights_dynamic = {
"early": {
"plot": {"core": 0.60, "scene": 0.20, "global": 0.20},
}
}
weights = manager._resolve_template_weights("plot", chapter=1)
assert weights == {"core": 0.60, "scene": 0.20, "global": 0.20}
def test_context_manager_genre_profile_fallbacks_to_project_info(temp_project):
manager = ContextManager(temp_project)
profile = manager._load_genre_profile({"project_info": {"genre": "xuanhuan"}})
assert profile.get("genre_raw") == "xuanhuan"
assert profile.get("genre") == "xuanhuan"
def test_context_manager_genre_profile_prefers_project_over_project_info(temp_project):
manager = ContextManager(temp_project)
profile = manager._load_genre_profile(
{
"project": {"genre": "xuanhuan"},
"project_info": {"genre": "dushi"},
}
)
assert profile.get("genre_raw") == "xuanhuan"
assert profile.get("genre") == "xuanhuan"
@@ -0,0 +1,55 @@
#!/usr/bin/env python3
# -*- coding: utf-8 -*-
from data_modules.config import DataModulesConfig
from data_modules.context_ranker import ContextRanker
def test_rank_recent_summaries_prefers_recency_and_hook(tmp_path):
cfg = DataModulesConfig.from_project_root(tmp_path)
ranker = ContextRanker(cfg)
items = [
{"chapter": 8, "summary": "平稳推进"},
{"chapter": 9, "summary": "最后留下悬念?"},
{"chapter": 7, "summary": "老信息"},
]
ranked = ranker.rank_recent_summaries(items, current_chapter=10)
assert ranked[0]["chapter"] == 9
assert ranked[-1]["chapter"] == 7
def test_rank_appearances_uses_recency_and_frequency(tmp_path):
cfg = DataModulesConfig.from_project_root(tmp_path)
ranker = ContextRanker(cfg)
items = [
{"entity_id": "a", "last_chapter": 9, "total": 1},
{"entity_id": "b", "last_chapter": 8, "total": 8},
{"entity_id": "c", "last_chapter": 9, "total": 3},
]
ranked = ranker.rank_appearances(items, current_chapter=10)
ids = [item["entity_id"] for item in ranked]
assert ids[0] == "c"
assert ids[-1] in {"a", "b"}
def test_rank_pack_adds_context_contract_meta(tmp_path):
cfg = DataModulesConfig.from_project_root(tmp_path)
ranker = ContextRanker(cfg)
pack = {
"meta": {"chapter": 12},
"core": {"recent_summaries": [{"chapter": 11, "summary": "x"}], "recent_meta": []},
"scene": {"appearing_characters": []},
"global": {},
"story_skeleton": [],
"alerts": {"disambiguation_warnings": [], "disambiguation_pending": []},
}
ranked = ranker.rank_pack(pack, chapter=12)
assert ranked["meta"]["context_contract_version"] == "v2"
assert ranked["meta"]["ranker"]["enabled"] is True
File diff suppressed because it is too large. Load diff
@@ -0,0 +1,97 @@
#!/usr/bin/env python3
# -*- coding: utf-8 -*-
"""
EntityLinker extra tests + CLI
"""
import sys
import pytest
from data_modules.entity_linker import EntityLinker, main as linker_main
from data_modules.index_manager import IndexManager, EntityMeta
@pytest.fixture
def temp_project(tmp_path):
from data_modules.config import DataModulesConfig
cfg = DataModulesConfig.from_project_root(tmp_path)
cfg.ensure_dirs()
return cfg
def test_process_extraction_and_register_new_entities(temp_project):
linker = EntityLinker(temp_project)
idx = IndexManager(temp_project)
idx.upsert_entity(
EntityMeta(
id="xiaoyan",
type="角色",
canonical_name="萧炎",
current={},
first_appearance=1,
last_appearance=1,
)
)
results, warnings = linker.process_extraction_result(
[
{
"mention": "萧炎",
"candidates": ["xiaoyan"],
"suggested": "xiaoyan",
"confidence": 0.7,
},
{
"mention": "宗主",
"candidates": ["zongzhu"],
"suggested": "zongzhu",
"confidence": 0.4,
},
]
)
assert len(results) == 2
assert len(warnings) == 2
registered = linker.register_new_entities(
[
{
"suggested_id": "hongyi",
"name": "红衣女子",
"type": "角色",
"mentions": ["红衣", "女子"],
}
]
)
assert registered == ["hongyi"]
aliases = idx.get_entity_aliases("hongyi")
assert "红衣女子" in aliases
def test_entity_linker_cli(temp_project, monkeypatch, capsys):
idx = IndexManager(temp_project)
idx.upsert_entity(
EntityMeta(
id="xiaoyan",
type="角色",
canonical_name="萧炎",
current={},
first_appearance=1,
last_appearance=1,
)
)
def run_cli(args):
monkeypatch.setattr(sys, "argv", ["entity_linker"] + args)
linker_main()
root = str(temp_project.project_root)
run_cli(["--project-root", root, "register-alias", "--entity", "xiaoyan", "--alias", "炎帝"])
run_cli(["--project-root", root, "lookup", "--mention", "炎帝"])
run_cli(["--project-root", root, "lookup", "--mention", "不存在"])
run_cli(["--project-root", root, "lookup-all", "--mention", "炎帝"])
run_cli(["--project-root", root, "list-aliases", "--entity", "xiaoyan"])
capsys.readouterr()
@@ -0,0 +1,273 @@
#!/usr/bin/env python3
# -*- coding: utf-8 -*-
import json
import sys
from pathlib import Path
def test_extract_state_summary_accepts_dominant_key(tmp_path):
scripts_dir = Path(__file__).resolve().parents[2]
if str(scripts_dir) not in sys.path:
sys.path.insert(0, str(scripts_dir))
from extract_chapter_context import extract_state_summary
state = {
"progress": {"current_chapter": 12, "total_words": 12345},
"protagonist_state": {
"power": {"realm": "筑基", "layer": 2},
"location": "宗门",
"golden_finger": {"name": "系统", "level": 1},
},
"strand_tracker": {
"history": [
{"chapter": 10, "dominant": "quest"},
{"chapter": 11, "dominant": "fire"},
]
},
}
noma_dir = tmp_path / ".noma"
noma_dir.mkdir(parents=True, exist_ok=True)
(noma_dir / "state.json").write_text(json.dumps(state, ensure_ascii=False), encoding="utf-8")
text = extract_state_summary(tmp_path)
assert "Ch10:quest" in text
assert "Ch11:fire" in text
def test_extract_chapter_outline_supports_hyphen_filename(tmp_path):
scripts_dir = Path(__file__).resolve().parents[2]
if str(scripts_dir) not in sys.path:
sys.path.insert(0, str(scripts_dir))
from extract_chapter_context import extract_chapter_outline
outline_dir = tmp_path / "大纲"
outline_dir.mkdir(parents=True, exist_ok=True)
(outline_dir / "第1卷-详细大纲.md").write_text("### 第1章:测试标题\n测试大纲", encoding="utf-8")
outline = extract_chapter_outline(tmp_path, 1)
assert "### 第1章:测试标题" in outline
assert "测试大纲" in outline
def test_extract_chapter_outline_prefers_state_volume_mapping(tmp_path):
scripts_dir = Path(__file__).resolve().parents[2]
if str(scripts_dir) not in sys.path:
sys.path.insert(0, str(scripts_dir))
from extract_chapter_context import extract_chapter_outline
noma_dir = tmp_path / ".noma"
noma_dir.mkdir(parents=True, exist_ok=True)
state = {
"progress": {
"volumes_planned": [
{"volume": 1, "chapters_range": "1-10"},
{"volume": 2, "chapters_range": "11-20"},
]
}
}
(noma_dir / "state.json").write_text(json.dumps(state, ensure_ascii=False), encoding="utf-8")
outline_dir = tmp_path / "大纲"
outline_dir.mkdir(parents=True, exist_ok=True)
(outline_dir / "第2卷-详细大纲.md").write_text("### 第12章:V2标题\nV2大纲", encoding="utf-8")
outline = extract_chapter_outline(tmp_path, 12)
assert "### 第12章:V2标题" in outline
assert "V2大纲" in outline
def test_extract_chapter_outline_falls_back_when_state_has_no_match(tmp_path):
scripts_dir = Path(__file__).resolve().parents[2]
if str(scripts_dir) not in sys.path:
sys.path.insert(0, str(scripts_dir))
from extract_chapter_context import extract_chapter_outline
noma_dir = tmp_path / ".noma"
noma_dir.mkdir(parents=True, exist_ok=True)
state = {"progress": {"volumes_planned": [{"volume": 1, "chapters_range": "1-10"}]}}
(noma_dir / "state.json").write_text(json.dumps(state, ensure_ascii=False), encoding="utf-8")
outline_dir = tmp_path / "大纲"
outline_dir.mkdir(parents=True, exist_ok=True)
(outline_dir / "第2卷-详细大纲.md").write_text("### 第60章:V2标题\nV2大纲", encoding="utf-8")
outline = extract_chapter_outline(tmp_path, 60)
assert "### 第60章:V2标题" in outline
assert "V2大纲" in outline
def test_build_chapter_context_payload_includes_contract_sections(tmp_path):
scripts_dir = Path(__file__).resolve().parents[2]
if str(scripts_dir) not in sys.path:
sys.path.insert(0, str(scripts_dir))
from extract_chapter_context import build_chapter_context_payload
from data_modules.config import DataModulesConfig
from data_modules.index_manager import IndexManager, ChapterReadingPowerMeta, ReviewMetrics
cfg = DataModulesConfig.from_project_root(tmp_path)
cfg.ensure_dirs()
state = {
"project": {"genre": "xuanhuan"},
"progress": {"current_chapter": 3, "total_words": 9000},
"protagonist_state": {
"power": {"realm": "筑基", "layer": 2},
"location": "宗门",
"golden_finger": {"name": "系统", "level": 1},
},
"strand_tracker": {"history": [{"chapter": 2, "dominant": "quest"}]},
"chapter_meta": {},
"disambiguation_warnings": [],
"disambiguation_pending": [],
}
(cfg.noma_dir / "state.json").write_text(json.dumps(state, ensure_ascii=False), encoding="utf-8")
summaries_dir = cfg.noma_dir / "summaries"
summaries_dir.mkdir(parents=True, exist_ok=True)
(summaries_dir / "ch0002.md").write_text("## 剧情摘要\n上一章总结", encoding="utf-8")
outline_dir = tmp_path / "大纲"
outline_dir.mkdir(parents=True, exist_ok=True)
(outline_dir / "第1卷 详细大纲.md").write_text("### 第3章:测试标题\n测试大纲", encoding="utf-8")
refs_dir = tmp_path / ".claude" / "references"
refs_dir.mkdir(parents=True, exist_ok=True)
(refs_dir / "genre-profiles.md").write_text("## xuanhuan\n- 升级线清晰", encoding="utf-8")
(refs_dir / "reading-power-taxonomy.md").write_text("## xuanhuan\n- 悬念钩优先", encoding="utf-8")
idx = IndexManager(cfg)
idx.save_chapter_reading_power(
ChapterReadingPowerMeta(chapter=2, hook_type="悬念钩", hook_strength="strong", coolpoint_patterns=["身份掉马"])
)
idx.save_review_metrics(
ReviewMetrics(start_chapter=1, end_chapter=2, overall_score=71, dimension_scores={"plot": 71})
)
payload = build_chapter_context_payload(tmp_path, 3)
assert payload["context_contract_version"] == "v2"
assert payload.get("context_weight_stage") in {"early", "mid", "late"}
assert "writing_guidance" in payload
assert isinstance(payload["writing_guidance"].get("guidance_items"), list)
assert isinstance(payload["writing_guidance"].get("checklist"), list)
assert isinstance(payload["writing_guidance"].get("checklist_score"), dict)
assert payload["genre_profile"].get("genre") == "xuanhuan"
assert "rag_assist" in payload
assert isinstance(payload["rag_assist"], dict)
assert payload["rag_assist"].get("invoked") is False
def test_render_text_contains_writing_guidance_section(tmp_path):
scripts_dir = Path(__file__).resolve().parents[2]
if str(scripts_dir) not in sys.path:
sys.path.insert(0, str(scripts_dir))
from extract_chapter_context import _render_text
payload = {
"chapter": 10,
"outline": "测试大纲",
"previous_summaries": ["### 第9章摘要\n上一章"],
"state_summary": "状态",
"context_contract_version": "v2",
"context_weight_stage": "early",
"reader_signal": {"review_trend": {"overall_avg": 72}, "low_score_ranges": [{"start_chapter": 8, "end_chapter": 9}]},
"genre_profile": {
"genre": "xuanhuan",
"genres": ["xuanhuan", "realistic"],
"composite_hints": ["以玄幻主线推进,同时保留现实议题表达"],
"reference_hints": ["升级线清晰"],
},
"writing_guidance": {
"guidance_items": ["先修低分", "钩子差异化"],
"checklist": [
{
"id": "fix_low_score_range",
"label": "修复低分区间问题",
"weight": 1.4,
"required": True,
"source": "reader_signal.low_score_ranges",
"verify_hint": "至少完成1处冲突升级",
}
],
"checklist_score": {
"score": 81.5,
"completion_rate": 0.66,
"required_completion_rate": 0.75,
},
"methodology": {
"enabled": True,
"framework": "digital-serial-v1",
"pilot": "xianxia",
"genre_profile_key": "xianxia",
"chapter_stage": "confront",
"observability": {
"next_reason_clarity": 78.0,
"anchor_effectiveness": 74.0,
"rhythm_naturalness": 72.0,
},
"signals": {"risk_flags": ["pattern_overuse_watch"]},
},
},
}
text = _render_text(payload)
assert "## 写作执行建议" in text
assert "先修低分" in text
assert "## Contract (v2)" in text
assert "- 上下文阶段权重: early" in text
assert "### 执行检查清单(可评分)" in text
assert "- 总权重: 1.40" in text
assert "[必做][w=1.4] 修复低分区间问题" in text
assert "### 执行评分" in text
assert "- 评分: 81.5" in text
assert "- 复合题材: xuanhuan + realistic" in text
assert "## 长篇方法论策略" in text
assert "- 适用题材: xianxia" in text
assert "next_reason=78.0" in text
def test_render_text_contains_rag_assist_section_when_hits_exist(tmp_path):
scripts_dir = Path(__file__).resolve().parents[2]
if str(scripts_dir) not in sys.path:
sys.path.insert(0, str(scripts_dir))
from extract_chapter_context import _render_text
payload = {
"chapter": 12,
"outline": "测试大纲",
"previous_summaries": [],
"state_summary": "状态",
"context_contract_version": "v2",
"reader_signal": {},
"genre_profile": {},
"writing_guidance": {},
"rag_assist": {
"invoked": True,
"mode": "auto",
"intent": "relationship",
"query": "第12章 人物关系与动机:萧炎与药老发生冲突",
"hits": [
{
"chapter": 9,
"scene_index": 2,
"source": "graph_hybrid",
"score": 0.91,
"content": "萧炎与药老在修炼方向上发生分歧。",
}
],
},
}
text = _render_text(payload)
assert "## RAG 检索线索" in text
assert "- 模式: auto" in text
assert "[graph_hybrid]" in text
assert "萧炎与药老" in text
@@ -0,0 +1,209 @@
#!/usr/bin/env python3
# -*- coding: utf-8 -*-
"""
migrate_state_to_sqlite tests
"""
import json
import pytest
import data_modules.migrate_state_to_sqlite as migrate_module
from data_modules.migrate_state_to_sqlite import (
migrate_state_to_sqlite,
_slim_world_settings,
_slim_relationships,
)
from data_modules.config import DataModulesConfig
from data_modules.index_manager import IndexManager
@pytest.fixture
def temp_project(tmp_path):
cfg = DataModulesConfig.from_project_root(tmp_path)
cfg.ensure_dirs()
return cfg
def test_migrate_state_missing_file(tmp_path):
cfg = DataModulesConfig.from_project_root(tmp_path)
stats = migrate_state_to_sqlite(cfg, dry_run=True, backup=False, verbose=False)
assert stats["entities"] == 0
def test_migrate_state_to_sqlite_flow(temp_project):
state = {
"entities_v3": {
"角色": {
"xiaoyan": {
"canonical_name": "萧炎",
"tier": "核心",
"desc": "主角",
"current": {"realm": "斗者"},
"first_appearance": 1,
"last_appearance": 2,
"is_protagonist": True,
}
}
},
"alias_index": {
"萧炎": [{"type": "角色", "id": "xiaoyan"}]
},
"state_changes": [
{"entity_id": "xiaoyan", "field": "realm", "old": "斗者", "new": "斗师", "reason": "突破", "chapter": 2}
],
"structured_relationships": [
{"from_entity": "xiaoyan", "to_entity": "yaolao", "type": "师徒", "description": "收徒", "chapter": 1}
],
"world_settings": {
"power_system": [{"name": "斗者"}, {"name": "斗师"}],
"factions": [{"name": "天云宗", "type": "宗门"}],
"locations": [{"name": "天云宗"}],
},
"plot_threads": {"active_threads": [], "foreshadowing": []},
"relationships": {},
"review_checkpoints": [],
"project_info": {"title": "测试书名"},
}
temp_project.state_file.write_text(json.dumps(state, ensure_ascii=False, indent=2), encoding="utf-8")
stats = migrate_state_to_sqlite(temp_project, dry_run=True, backup=False, verbose=False)
assert stats["entities"] == 1
assert stats["aliases"] == 1
stats = migrate_state_to_sqlite(temp_project, dry_run=False, backup=False, verbose=False)
assert stats["entities"] == 1
# state.json 被精简
saved = json.loads(temp_project.state_file.read_text(encoding="utf-8"))
assert saved.get("_migrated_to_sqlite") is True
assert "entities_v3" not in saved
# SQLite 中可查询实体
idx = IndexManager(temp_project)
entity = idx.get_entity("xiaoyan")
assert entity is not None
def test_slim_helpers():
world = {
"power_system": [{"name": "斗者"}],
"factions": [{"name": "天云宗", "type": "宗门"}],
"locations": [{"name": "天云宗"}],
}
slim = _slim_world_settings(world)
assert slim["power_system"][0] == "斗者"
rels = _slim_relationships({"a": 1})
assert rels["a"] == 1
def test_slim_helpers_non_dict():
assert _slim_world_settings("bad") == {}
assert _slim_relationships("bad") == {}
def test_migrate_state_verbose_and_dry_run(temp_project, capsys):
state = {
"entities_v3": {},
"alias_index": {},
"state_changes": [],
"structured_relationships": [],
"world_settings": {},
"plot_threads": {},
"relationships": {},
"review_checkpoints": [],
"project_info": {},
}
temp_project.state_file.write_text(json.dumps(state, ensure_ascii=False), encoding="utf-8")
stats = migrate_state_to_sqlite(temp_project, dry_run=True, backup=False, verbose=True)
output = capsys.readouterr().out
assert stats["errors"] == 0
assert "dry-run" in output or "dry run" in output
def test_migrate_state_cli_main(tmp_path, monkeypatch, capsys):
project_root = tmp_path
args = [
"migrate_state_to_sqlite",
"--project-root",
str(project_root),
"--dry-run",
"--no-backup",
]
monkeypatch.setattr("sys.argv", args)
migrate_module.main()
output = json.loads(capsys.readouterr().out or "{}")
assert output.get("status") == "success"
def test_migrate_state_backup_and_skips(temp_project):
state = {
"entities_v3": {
"角色": {
"good": {"canonical_name": "好人"},
"bad": "not-dict",
}
},
"alias_index": {
"好人": [{"type": "角色", "id": "good"}],
"坏条目": ["oops", {"type": "角色"}],
},
"state_changes": ["bad", {"field": "realm"}],
"structured_relationships": ["bad", {"from_entity": "", "to_entity": ""}],
"relationships": {},
"world_settings": {},
"plot_threads": {},
"review_checkpoints": [],
"project_info": {},
}
temp_project.state_file.write_text(json.dumps(state, ensure_ascii=False), encoding="utf-8")
stats = migrate_state_to_sqlite(temp_project, dry_run=False, backup=True, verbose=False)
assert stats["entities"] == 1
assert stats["skipped"] >= 3
backups = list(temp_project.state_file.parent.glob("state.json.backup-*"))
assert backups
def test_migrate_state_error_branches(tmp_path, monkeypatch):
cfg = DataModulesConfig.from_project_root(tmp_path)
cfg.ensure_dirs()
state = {
"entities_v3": {"角色": {"boom": {"canonical_name": "爆"}}},
"alias_index": {"爆": [{"type": "角色", "id": "boom"}]},
"state_changes": [
{"entity_id": "boom", "field": "realm", "old": "", "new": "斗者", "reason": "测试", "chapter": 1}
],
"structured_relationships": [
{"from_entity": "boom", "to_entity": "yao", "type": "相识", "description": "测试", "chapter": 1}
],
"relationships": {},
"world_settings": {},
"plot_threads": {},
"review_checkpoints": [],
"project_info": {},
}
cfg.state_file.write_text(json.dumps(state, ensure_ascii=False), encoding="utf-8")
class BoomSQL:
def __init__(self, *args, **kwargs):
pass
def upsert_entity(self, *args, **kwargs):
raise RuntimeError("boom")
def register_alias(self, *args, **kwargs):
raise RuntimeError("boom")
def record_state_change(self, *args, **kwargs):
raise RuntimeError("boom")
def upsert_relationship(self, *args, **kwargs):
raise RuntimeError("boom")
monkeypatch.setattr(migrate_module, "SQLStateManager", BoomSQL)
stats = migrate_state_to_sqlite(cfg, dry_run=False, backup=False, verbose=False)
assert stats["errors"] >= 4
@@ -0,0 +1,106 @@
#!/usr/bin/env python3
# -*- coding: utf-8 -*-
import sys
from pathlib import Path
def _ensure_scripts_on_path() -> None:
scripts_dir = Path(__file__).resolve().parents[2]
if str(scripts_dir) not in sys.path:
sys.path.insert(0, str(scripts_dir))
def test_resolve_project_root_prefers_cwd_project(tmp_path):
_ensure_scripts_on_path()
from project_locator import resolve_project_root
project_root = tmp_path / "workspace"
(project_root / ".noma").mkdir(parents=True, exist_ok=True)
(project_root / ".noma" / "state.json").write_text("{}", encoding="utf-8")
resolved = resolve_project_root(cwd=project_root)
assert resolved == project_root.resolve()
def test_resolve_project_root_stops_at_git_root(tmp_path):
_ensure_scripts_on_path()
from project_locator import resolve_project_root
repo_root = tmp_path / "repo"
(repo_root / ".git").mkdir(parents=True, exist_ok=True)
nested = repo_root / "sub" / "dir"
nested.mkdir(parents=True, exist_ok=True)
outside_project = tmp_path / "outside_project"
(outside_project / ".noma").mkdir(parents=True, exist_ok=True)
(outside_project / ".noma" / "state.json").write_text("{}", encoding="utf-8")
try:
resolve_project_root(cwd=nested)
assert False, "Expected FileNotFoundError when only parent outside git root has project"
except FileNotFoundError:
pass
def test_resolve_project_root_finds_default_subdir_within_git_root(tmp_path):
_ensure_scripts_on_path()
from project_locator import resolve_project_root
repo_root = tmp_path / "repo"
(repo_root / ".git").mkdir(parents=True, exist_ok=True)
default_project = repo_root / "noma-project"
(default_project / ".noma").mkdir(parents=True, exist_ok=True)
(default_project / ".noma" / "state.json").write_text("{}", encoding="utf-8")
nested = repo_root / "sub" / "dir"
nested.mkdir(parents=True, exist_ok=True)
resolved = resolve_project_root(cwd=nested)
assert resolved == default_project.resolve()
def test_resolve_project_root_uses_workspace_pointer(tmp_path):
_ensure_scripts_on_path()
from project_locator import resolve_project_root, write_current_project_pointer
workspace = tmp_path / "workspace"
(workspace / ".claude").mkdir(parents=True, exist_ok=True)
project_root = workspace / "凡人资本论"
(project_root / ".noma").mkdir(parents=True, exist_ok=True)
(project_root / ".noma" / "state.json").write_text("{}", encoding="utf-8")
pointer_file = write_current_project_pointer(project_root, workspace_root=workspace)
assert pointer_file is not None
assert pointer_file.is_file()
resolved = resolve_project_root(cwd=workspace)
assert resolved == project_root.resolve()
def test_resolve_project_root_ignores_stale_pointer_and_fallbacks(tmp_path):
_ensure_scripts_on_path()
from project_locator import resolve_project_root
workspace = tmp_path / "workspace"
(workspace / ".claude").mkdir(parents=True, exist_ok=True)
# stale pointer
(workspace / ".claude" / ".noma-current-project").write_text(
str(workspace / "missing-project"), encoding="utf-8"
)
default_project = workspace / "noma-project"
(default_project / ".noma").mkdir(parents=True, exist_ok=True)
(default_project / ".noma" / "state.json").write_text("{}", encoding="utf-8")
resolved = resolve_project_root(cwd=workspace)
assert resolved == default_project.resolve()
@@ -0,0 +1,513 @@
#!/usr/bin/env python3
# -*- coding: utf-8 -*-
"""
RAGAdapter tests
"""
import sys
import json
import asyncio
import logging
import sqlite3
from contextlib import closing
import pytest
import data_modules.rag_adapter as rag_module
from data_modules.rag_adapter import RAGAdapter
from data_modules.config import DataModulesConfig
from data_modules.index_manager import EntityMeta, RelationshipMeta
class StubClient:
async def embed(self, texts):
return [[1.0, 0.0] for _ in texts]
async def embed_batch(self, texts, skip_failures=True):
return [[1.0, 0.0] for _ in texts]
async def rerank(self, query, documents, top_n=None):
top_n = top_n or len(documents)
return [{"index": i, "relevance_score": 1.0 / (i + 1)} for i in range(min(top_n, len(documents)))]
class StubClientWithFailures(StubClient):
async def embed_batch(self, texts, skip_failures=True):
if len(texts) == 1:
return [None]
return [None, [1.0, 0.0]]
class StubEmbedClient401:
def __init__(self):
self.last_error_status = 401
self.last_error_message = "auth failed"
class StubClientAuthFailure(StubClient):
def __init__(self):
self._embed_client = StubEmbedClient401()
async def embed(self, texts):
return None
class StubClientRerankFailure(StubClient):
async def rerank(self, query, documents, top_n=None):
return []
@pytest.fixture
def temp_project(tmp_path, monkeypatch):
cfg = DataModulesConfig.from_project_root(tmp_path)
cfg.ensure_dirs()
monkeypatch.setattr(rag_module, "get_client", lambda config: StubClient())
return cfg
@pytest.mark.asyncio
async def test_store_and_search(temp_project):
adapter = RAGAdapter(temp_project)
chunks = [
{"chapter": 1, "scene_index": 1, "content": "萧炎在天云宗修炼斗气"},
{"chapter": 1, "scene_index": 2, "content": "药老传授炼药技巧"},
]
stored = await adapter.store_chunks(chunks)
assert stored == 2
vec_results = await adapter.vector_search("萧炎", top_k=2)
assert len(vec_results) == 2
bm25_results = adapter.bm25_search("萧炎", top_k=2)
assert len(bm25_results) >= 1
stats = adapter.get_stats()
assert stats["vectors"] == 2
@pytest.mark.asyncio
async def test_store_chunks_with_embedding_failure(tmp_path, monkeypatch):
cfg = DataModulesConfig.from_project_root(tmp_path)
cfg.ensure_dirs()
monkeypatch.setattr(rag_module, "get_client", lambda config: StubClientWithFailures())
adapter = RAGAdapter(cfg)
chunks = [
{"chapter": 1, "scene_index": 1, "content": "短内容"},
{"chapter": 1, "scene_index": 2, "content": "稍长内容用于索引"},
]
stored = await adapter.store_chunks(chunks)
assert stored == 1
@pytest.mark.asyncio
async def test_hybrid_search_full_scan(temp_project):
adapter = RAGAdapter(temp_project)
await adapter.store_chunks(
[{"chapter": 1, "scene_index": 1, "content": "萧炎修炼"}]
)
results = await adapter.hybrid_search("萧炎", vector_top_k=5, bm25_top_k=5, rerank_top_n=1)
assert results
assert results[0].source == "hybrid"
@pytest.mark.asyncio
async def test_hybrid_search_prefilter(tmp_path, monkeypatch):
cfg = DataModulesConfig.from_project_root(tmp_path)
cfg.ensure_dirs()
cfg.vector_full_scan_max_vectors = 0
monkeypatch.setattr(rag_module, "get_client", lambda config: StubClient())
adapter = RAGAdapter(cfg)
await adapter.store_chunks(
[
{"chapter": 1, "scene_index": 1, "content": "萧炎修炼"},
{"chapter": 2, "scene_index": 1, "content": "药老出场"},
]
)
results = await adapter.hybrid_search("药老", vector_top_k=2, bm25_top_k=2, rerank_top_n=1)
assert results
@pytest.mark.asyncio
async def test_search_respects_chapter_filter_across_strategies(tmp_path, monkeypatch):
cfg = DataModulesConfig.from_project_root(tmp_path)
cfg.ensure_dirs()
cfg.vector_full_scan_max_vectors = 0 # 强制走预筛选分支
monkeypatch.setattr(rag_module, "get_client", lambda config: StubClient())
adapter = RAGAdapter(cfg)
await adapter.store_chunks(
[
{"chapter": 1, "scene_index": 1, "content": "前文线索,尚未涉及关键宝物"},
{"chapter": 2, "scene_index": 1, "content": "秘宝现世,引发争夺"},
{"chapter": 3, "scene_index": 1, "content": "秘宝大战彻底爆发"},
]
)
vector_results = await adapter.vector_search("秘宝", top_k=5, chapter=1)
assert vector_results
assert all((r.chapter or 0) <= 1 for r in vector_results)
bm25_results = adapter.bm25_search("秘宝", top_k=5, chapter=1)
assert bm25_results
assert all((r.chapter or 0) <= 1 for r in bm25_results)
hybrid_results = await adapter.hybrid_search(
"秘宝",
vector_top_k=5,
bm25_top_k=5,
rerank_top_n=3,
chapter=1,
)
assert hybrid_results
assert all((r.chapter or 0) <= 1 for r in hybrid_results)
@pytest.mark.asyncio
async def test_graph_hybrid_search_with_entity_expansion(tmp_path, monkeypatch):
cfg = DataModulesConfig.from_project_root(tmp_path)
cfg.ensure_dirs()
cfg.graph_rag_enabled = True
monkeypatch.setattr(rag_module, "get_client", lambda config: StubClient())
adapter = RAGAdapter(cfg)
adapter.index_manager.upsert_entity(
EntityMeta(
id="xiaoyan",
type="角色",
canonical_name="萧炎",
current={},
first_appearance=1,
last_appearance=2,
)
)
adapter.index_manager.upsert_entity(
EntityMeta(
id="yaolao",
type="角色",
canonical_name="药老",
current={},
first_appearance=1,
last_appearance=2,
)
)
adapter.index_manager.register_alias("萧炎", "xiaoyan", "角色")
adapter.index_manager.register_alias("药老", "yaolao", "角色")
adapter.index_manager.upsert_relationship(
RelationshipMeta(
from_entity="xiaoyan",
to_entity="yaolao",
type="师徒",
description="收徒",
chapter=1,
)
)
await adapter.store_chunks(
[
{"chapter": 1, "scene_index": 1, "content": "萧炎拜药老为师,正式成为师徒"},
{"chapter": 2, "scene_index": 1, "content": "萧炎在天云宗修炼斗气"},
]
)
results = await adapter.graph_hybrid_search(
"萧炎和药老关系",
top_k=2,
center_entities=["萧炎", "药老"],
)
assert results
assert any("药老" in r.content for r in results)
assert all(r.source == "graph_hybrid" for r in results)
@pytest.mark.asyncio
async def test_search_auto_uses_graph_strategy_when_enabled(tmp_path, monkeypatch):
cfg = DataModulesConfig.from_project_root(tmp_path)
cfg.ensure_dirs()
cfg.graph_rag_enabled = True
monkeypatch.setattr(rag_module, "get_client", lambda config: StubClient())
adapter = RAGAdapter(cfg)
adapter.index_manager.upsert_entity(
EntityMeta(
id="xiaoyan",
type="角色",
canonical_name="萧炎",
current={},
first_appearance=1,
last_appearance=1,
)
)
adapter.index_manager.register_alias("萧炎", "xiaoyan", "角色")
await adapter.store_chunks(
[{"chapter": 1, "scene_index": 1, "content": "萧炎突破斗师"}]
)
results = await adapter.search("萧炎关系", top_k=1, strategy="auto")
assert results
assert results[0].source in {"graph_hybrid", "hybrid"}
@pytest.mark.asyncio
async def test_graph_hybrid_search_fallback_when_graph_disabled(tmp_path, monkeypatch):
cfg = DataModulesConfig.from_project_root(tmp_path)
cfg.ensure_dirs()
cfg.graph_rag_enabled = False
monkeypatch.setattr(rag_module, "get_client", lambda config: StubClient())
adapter = RAGAdapter(cfg)
await adapter.store_chunks(
[{"chapter": 1, "scene_index": 1, "content": "萧炎在天云宗修炼斗气"}]
)
modes = []
def _record_log(query, mode, results, latency_ms, chapter=None):
modes.append(mode)
monkeypatch.setattr(adapter, "_log_query", _record_log)
results = await adapter.graph_hybrid_search("萧炎关系", top_k=1)
assert results
assert modes
assert modes[-1] == "graph_hybrid_fallback"
assert all(r.source == "hybrid" for r in results)
@pytest.mark.asyncio
async def test_graph_hybrid_search_rerank_failure_uses_candidates(tmp_path, monkeypatch):
cfg = DataModulesConfig.from_project_root(tmp_path)
cfg.ensure_dirs()
cfg.graph_rag_enabled = True
monkeypatch.setattr(rag_module, "get_client", lambda config: StubClientRerankFailure())
adapter = RAGAdapter(cfg)
adapter.index_manager.upsert_entity(
EntityMeta(
id="xiaoyan",
type="角色",
canonical_name="萧炎",
current={},
first_appearance=1,
last_appearance=2,
)
)
adapter.index_manager.upsert_entity(
EntityMeta(
id="yaolao",
type="角色",
canonical_name="药老",
current={},
first_appearance=1,
last_appearance=2,
)
)
adapter.index_manager.register_alias("萧炎", "xiaoyan", "角色")
adapter.index_manager.register_alias("药老", "yaolao", "角色")
adapter.index_manager.upsert_relationship(
RelationshipMeta(
from_entity="xiaoyan",
to_entity="yaolao",
type="师徒",
description="收徒",
chapter=1,
)
)
await adapter.store_chunks(
[
{"chapter": 1, "scene_index": 1, "content": "萧炎拜药老为师,正式成为师徒"},
{"chapter": 2, "scene_index": 1, "content": "萧炎在天云宗修炼斗气"},
]
)
results = await adapter.graph_hybrid_search(
"萧炎和药老关系",
top_k=2,
center_entities=["萧炎", "药老"],
)
assert results
assert len(results) <= 2
assert all(r.source == "graph_hybrid" for r in results)
@pytest.mark.asyncio
async def test_search_unknown_strategy_falls_back_to_hybrid(tmp_path, monkeypatch):
cfg = DataModulesConfig.from_project_root(tmp_path)
cfg.ensure_dirs()
monkeypatch.setattr(rag_module, "get_client", lambda config: StubClient())
adapter = RAGAdapter(cfg)
await adapter.store_chunks(
[{"chapter": 1, "scene_index": 1, "content": "萧炎在天云宗修炼斗气"}]
)
results = await adapter.search("萧炎", top_k=1, strategy="not_exists")
assert results
assert all(r.source == "hybrid" for r in results)
@pytest.mark.asyncio
async def test_search_with_backtrack(temp_project):
adapter = RAGAdapter(temp_project)
chunks = [
{
"chapter": 1,
"scene_index": 0,
"content": "章节摘要",
"chunk_type": "summary",
"chunk_id": "ch0001_summary",
"source_file": "summaries/ch0001.md",
},
{
"chapter": 1,
"scene_index": 1,
"content": "场景内容",
"chunk_type": "scene",
"chunk_id": "ch0001_s1",
"parent_chunk_id": "ch0001_summary",
"source_file": "正文/第0001章.md#scene_1",
},
]
await adapter.store_chunks(chunks)
results = await adapter.search_with_backtrack("场景", top_k=1)
assert any(r.chunk_type == "summary" for r in results)
def test_vector_helpers(temp_project):
adapter = RAGAdapter(temp_project)
emb = [1.0, 0.0]
data = adapter._serialize_embedding(emb)
assert adapter._deserialize_embedding(data) == emb
assert adapter._cosine_similarity([0.0, 0.0], [1.0, 0.0]) == 0.0
def test_recent_and_fetch_vectors(temp_project):
adapter = RAGAdapter(temp_project)
with adapter._get_conn() as conn:
cursor = conn.cursor()
cursor.execute(
"INSERT INTO vectors (chunk_id, chapter, scene_index, content, embedding, parent_chunk_id, chunk_type, source_file) VALUES (?, ?, ?, ?, ?, ?, ?, ?)",
("ch0001_s1", 1, 1, "内容", b"", None, "scene", "正文/第0001章.md#scene_1"),
)
cursor.execute(
"INSERT INTO vectors (chunk_id, chapter, scene_index, content, embedding, parent_chunk_id, chunk_type, source_file) VALUES (?, ?, ?, ?, ?, ?, ?, ?)",
("ch0002_s1", 2, 1, "后文内容", b"", None, "scene", "正文/第0002章.md#scene_1"),
)
conn.commit()
assert adapter._get_vectors_count() == 2
assert adapter._get_recent_chunk_ids(1) == ["ch0002_s1"]
assert adapter._get_recent_chunk_ids(10, chapter=1) == ["ch0001_s1"]
rows = adapter._fetch_vectors_by_chunk_ids(["ch0001_s1"])
assert len(rows) == 1
def test_init_db_migrates_legacy_vectors_schema(tmp_path, monkeypatch):
cfg = DataModulesConfig.from_project_root(tmp_path)
cfg.ensure_dirs()
monkeypatch.setattr(rag_module, "get_client", lambda config: StubClient())
# 旧结构:缺少 parent_chunk_id/chunk_type/source_file/created_at
with closing(sqlite3.connect(str(cfg.vector_db))) as conn:
cursor = conn.cursor()
cursor.execute(
"""
CREATE TABLE vectors (
chunk_id TEXT PRIMARY KEY,
chapter INTEGER,
scene_index INTEGER,
content TEXT,
embedding BLOB
)
"""
)
cursor.execute(
"""
INSERT INTO vectors (chunk_id, chapter, scene_index, content, embedding)
VALUES (?, ?, ?, ?, ?)
""",
("ch0001_s1", 1, 1, "旧数据", b""),
)
conn.commit()
adapter = RAGAdapter(cfg)
with adapter._get_conn() as conn:
cursor = conn.cursor()
cursor.execute("PRAGMA table_info(vectors)")
cols = {row[1] for row in cursor.fetchall()}
assert {"parent_chunk_id", "chunk_type", "source_file", "created_at"}.issubset(cols)
cursor.execute("SELECT COUNT(*) FROM vectors")
assert cursor.fetchone()[0] == 1
cursor.execute("SELECT chunk_type FROM vectors WHERE chunk_id = ?", ("ch0001_s1",))
row = cursor.fetchone()
assert row is not None
assert row[0] == "scene"
backup_dir = cfg.noma_dir / "backups"
backups = list(backup_dir.glob("vectors.db.schema_migration.v*.bak"))
assert backups
def test_rag_adapter_cli(temp_project, monkeypatch, capsys):
# stats
def run_cli(args):
monkeypatch.setattr(sys, "argv", ["rag_adapter"] + args)
rag_module.main()
root = str(temp_project.project_root)
run_cli(["--project-root", root, "stats"])
# index-chapter
run_cli(
[
"--project-root",
root,
"index-chapter",
"--chapter",
"1",
"--scenes",
json.dumps([{"index": 1, "summary": "摘要", "content": "内容"}], ensure_ascii=False),
]
)
# search
run_cli(["--project-root", root, "search", "--query", "内容", "--mode", "bm25", "--top-k", "5"])
run_cli(["--project-root", root, "search", "--query", "内容", "--mode", "vector", "--top-k", "5"])
run_cli(["--project-root", root, "search", "--query", "内容", "--mode", "hybrid", "--top-k", "5"])
run_cli(["--project-root", root, "search", "--query", "内容", "--mode", "auto", "--top-k", "5"])
capsys.readouterr()
def test_rag_adapter_log_query_failure_is_reported(temp_project, monkeypatch, caplog):
adapter = RAGAdapter(temp_project)
def _raise_log_error(*args, **kwargs):
raise RuntimeError("log write failed")
monkeypatch.setattr(adapter.index_manager, "log_rag_query", _raise_log_error)
with caplog.at_level(logging.WARNING):
adapter._log_query("q", "vector", [], 1)
message_text = "\n".join(record.getMessage() for record in caplog.records)
assert "failed to log rag query" in message_text
def test_rag_adapter_cli_search_shows_degraded_warning(temp_project, monkeypatch, capsys):
monkeypatch.setattr(rag_module, "get_client", lambda config: StubClientAuthFailure())
def run_cli(args):
monkeypatch.setattr(sys, "argv", ["rag_adapter"] + args)
rag_module.main()
root = str(temp_project.project_root)
run_cli(["--project-root", root, "search", "--query", "测试", "--mode", "vector", "--top-k", "3"])
captured = capsys.readouterr()
payload = json.loads(captured.out.strip().splitlines()[-1])
assert payload.get("status") == "success"
warnings = payload.get("warnings") or []
assert warnings
assert warnings[0].get("code") == "DEGRADED_MODE"
assert warnings[0].get("reason") == "embedding_auth_failed"
@@ -0,0 +1,333 @@
#!/usr/bin/env python3
# -*- coding: utf-8 -*-
"""
关系事件与关系图谱测试
"""
import json
import sys
import pytest
import data_modules.index_manager as index_manager_module
from data_modules.config import DataModulesConfig
from data_modules.index_manager import (
EntityMeta,
IndexManager,
RelationshipEventMeta,
RelationshipMeta,
)
@pytest.fixture
def temp_project(tmp_path):
cfg = DataModulesConfig.from_project_root(tmp_path)
cfg.ensure_dirs()
return cfg
def test_relationship_events_timeline_and_subgraph(temp_project):
manager = IndexManager(temp_project)
manager.upsert_entity(
EntityMeta(
id="xiaoyan",
type="角色",
canonical_name="萧炎",
tier="核心",
current={},
first_appearance=1,
last_appearance=10,
is_protagonist=True,
)
)
manager.upsert_entity(
EntityMeta(
id="yaolao",
type="角色",
canonical_name="药老",
tier="重要",
current={},
first_appearance=1,
last_appearance=10,
)
)
manager.upsert_entity(
EntityMeta(
id="lintian",
type="角色",
canonical_name="林天",
tier="重要",
current={},
first_appearance=2,
last_appearance=10,
)
)
manager.upsert_relationship(
RelationshipMeta(
from_entity="xiaoyan",
to_entity="yaolao",
type="师徒",
description="正式拜师",
chapter=3,
)
)
manager.upsert_relationship(
RelationshipMeta(
from_entity="yaolao",
to_entity="lintian",
type="敌对",
description="理念冲突",
chapter=5,
)
)
event_id = manager.record_relationship_event(
RelationshipEventMeta(
from_entity="xiaoyan",
to_entity="yaolao",
type="师徒",
chapter=3,
action="create",
polarity=1,
strength=0.9,
description="拜师",
evidence="公开收徒",
confidence=0.95,
)
)
assert event_id > 0
manager.record_relationship_event(
RelationshipEventMeta(
from_entity="yaolao",
to_entity="lintian",
type="敌对",
chapter=5,
action="create",
polarity=-1,
strength=0.8,
description="结怨",
evidence="比斗失手",
confidence=0.8,
)
)
events = manager.get_relationship_events("xiaoyan", direction="both", limit=20)
assert events
timeline = manager.get_relationship_timeline("xiaoyan", "yaolao", limit=20)
assert timeline
assert timeline[0]["type"] == "师徒"
graph = manager.build_relationship_subgraph("xiaoyan", depth=2, chapter=10, top_edges=10)
node_ids = {n["id"] for n in graph["nodes"]}
assert "xiaoyan" in node_ids
assert "yaolao" in node_ids
assert "lintian" in node_ids
assert graph["edges"]
mermaid = manager.render_relationship_subgraph_mermaid(graph)
assert "mermaid" in mermaid
assert "师徒" in mermaid
def test_relationship_subgraph_respects_chapter_slice(temp_project):
manager = IndexManager(temp_project)
manager.upsert_entity(
EntityMeta(
id="a",
type="角色",
canonical_name="甲",
current={},
first_appearance=1,
last_appearance=3,
is_protagonist=True,
)
)
manager.upsert_entity(
EntityMeta(
id="b",
type="角色",
canonical_name="乙",
current={},
first_appearance=1,
last_appearance=3,
)
)
manager.record_relationship_event(
RelationshipEventMeta(
from_entity="a",
to_entity="b",
type="同盟",
chapter=1,
action="create",
polarity=1,
strength=0.6,
)
)
manager.record_relationship_event(
RelationshipEventMeta(
from_entity="a",
to_entity="b",
type="同盟",
chapter=2,
action="remove",
polarity=0,
strength=0.0,
)
)
graph_ch1 = manager.build_relationship_subgraph("a", depth=1, chapter=1, top_edges=10)
graph_ch3 = manager.build_relationship_subgraph("a", depth=1, chapter=3, top_edges=10)
assert len(graph_ch1["edges"]) == 1
assert len(graph_ch3["edges"]) == 0
def test_relationship_subgraph_fallbacks_to_snapshot_when_events_missing(temp_project):
manager = IndexManager(temp_project)
manager.upsert_entity(
EntityMeta(
id="a",
type="角色",
canonical_name="甲",
current={},
first_appearance=1,
last_appearance=5,
is_protagonist=True,
)
)
manager.upsert_entity(
EntityMeta(
id="b",
type="角色",
canonical_name="乙",
current={},
first_appearance=1,
last_appearance=5,
)
)
# 只写 relationships 快照,不写 relationship_events
manager.upsert_relationship(
RelationshipMeta(
from_entity="a",
to_entity="b",
type="同盟",
description="旧版快照数据",
chapter=3,
)
)
graph = manager.build_relationship_subgraph("a", depth=1, chapter=3, top_edges=10)
assert graph["edges"]
assert graph["edges"][0]["action"] == "snapshot"
assert graph["edges"][0]["type"] == "同盟"
def test_relationship_graph_cli_commands(temp_project, monkeypatch, capsys):
manager = IndexManager(temp_project)
manager.upsert_entity(
EntityMeta(
id="hero",
type="角色",
canonical_name="主角",
current={},
first_appearance=1,
last_appearance=1,
is_protagonist=True,
)
)
manager.upsert_entity(
EntityMeta(
id="mentor",
type="角色",
canonical_name="师父",
current={},
first_appearance=1,
last_appearance=1,
)
)
manager.record_relationship_event(
RelationshipEventMeta(
from_entity="hero",
to_entity="mentor",
type="师徒",
chapter=1,
action="create",
polarity=1,
strength=0.9,
)
)
root = str(temp_project.project_root)
def run_cli(args):
monkeypatch.setattr(sys, "argv", ["index_manager"] + args)
index_manager_module.main()
output = capsys.readouterr().out.strip().splitlines()
assert output
return json.loads(output[-1])
payload = run_cli(
[
"--project-root",
root,
"get-relationship-events",
"--entity",
"hero",
"--direction",
"both",
"--limit",
"10",
]
)
assert payload["status"] == "success"
assert payload["data"]
payload = run_cli(
[
"--project-root",
root,
"get-relationship-graph",
"--center",
"hero",
"--depth",
"1",
"--chapter",
"1",
"--format",
"mermaid",
]
)
assert payload["status"] == "success"
assert "mermaid" in payload["data"]["mermaid"]
payload = run_cli(
[
"--project-root",
root,
"get-relationship-timeline",
"--a",
"hero",
"--b",
"mentor",
"--limit",
"10",
]
)
assert payload["status"] == "success"
assert payload["data"]
payload = run_cli(
[
"--project-root",
root,
"record-relationship-event",
"--data",
json.dumps(
{
"from_entity": "hero",
"type": "师徒",
"chapter": 1,
},
ensure_ascii=False,
),
]
)
assert payload["status"] == "error"
assert payload["error"]["code"] == "INVALID_RELATIONSHIP_EVENT"
@@ -0,0 +1,213 @@
#!/usr/bin/env python3
# -*- coding: utf-8 -*-
"""
SQLStateManager tests
"""
import json
import sys
import pytest
import data_modules.sql_state_manager as sql_state_manager_module
from data_modules.sql_state_manager import SQLStateManager, EntityData
from data_modules.index_manager import EntityMeta
@pytest.fixture
def temp_project(tmp_path):
from data_modules.config import DataModulesConfig
cfg = DataModulesConfig.from_project_root(tmp_path)
cfg.ensure_dirs()
return cfg
def test_sql_state_manager_entity_and_alias(temp_project):
manager = SQLStateManager(temp_project)
entity = EntityData(
id="xiaoyan",
type="角色",
name="萧炎",
tier="核心",
current={"realm": "斗师"},
aliases=["炎帝", "小炎子"],
is_protagonist=True,
)
assert manager.upsert_entity(entity) is True
assert manager.upsert_entity(entity) is False
fetched = manager.get_entity("xiaoyan")
assert "炎帝" in fetched["aliases"]
by_type = manager.get_entities_by_type("角色")
assert any(e["id"] == "xiaoyan" for e in by_type)
core = manager.get_core_entities()
assert any(e["id"] == "xiaoyan" for e in core)
protagonist = manager.get_protagonist()
assert protagonist["id"] == "xiaoyan"
resolved = manager.resolve_alias("炎帝")
assert any(r["id"] == "xiaoyan" for r in resolved)
assert manager.update_entity_current("xiaoyan", {"realm": "斗王"}) is True
updated = manager.get_entity("xiaoyan")
assert updated["current_json"]["realm"] == "斗王"
def test_sql_state_manager_state_changes_and_relationships(temp_project):
manager = SQLStateManager(temp_project)
manager.upsert_entity(
EntityData(id="xiaoyan", type="角色", name="萧炎", current={})
)
change_id = manager.record_state_change(
entity_id="xiaoyan",
field="realm",
old_value="斗者",
new_value="斗师",
reason="突破",
chapter=2,
)
assert change_id > 0
assert len(manager.get_entity_state_changes("xiaoyan")) == 1
assert len(manager.get_recent_state_changes(limit=5)) == 1
assert len(manager.get_chapter_state_changes(2)) == 1
assert manager.upsert_relationship(
from_entity="xiaoyan",
to_entity="yaolao",
type="师徒",
description="收徒",
chapter=1,
)
rels = manager.get_entity_relationships("xiaoyan", direction="from")
assert len(rels) == 1
between = manager.get_relationship_between("xiaoyan", "yaolao")
assert len(between) == 1
assert len(manager.get_recent_relationships(limit=5)) >= 1
def test_sql_state_manager_process_chapter_entities_and_exports(temp_project):
manager = SQLStateManager(temp_project)
stats = manager.process_chapter_entities(
chapter=10,
entities_appeared=[{"id": "xiaoyan", "mentions": ["萧炎"], "confidence": 0.9}],
entities_new=[
{"suggested_id": "yaolao", "name": "药老", "type": "角色", "tier": "重要"}
],
state_changes=[
{"entity_id": "yaolao", "field": "status", "old": "", "new": "出场", "reason": "登场"}
],
relationships_new=[
{"from": "xiaoyan", "to": "yaolao", "type": "师徒", "description": "收徒"}
],
)
assert stats["entities_created"] >= 1
assert stats["relationships"] == 1
rel_events = manager._index_manager.get_relationship_events("xiaoyan", direction="both")
assert len(rel_events) >= 1
entities_v3 = manager.export_to_entities_v3_format()
assert "角色" in entities_v3
alias_index = manager.export_to_alias_index_format()
assert isinstance(alias_index, dict)
def test_sql_state_manager_existing_entity_updates_and_stats(temp_project):
manager = SQLStateManager(temp_project)
manager.upsert_entity(
EntityData(id="xiaoyan", type="角色", name="萧炎", current={"hp": 5})
)
stats = manager.process_chapter_entities(
chapter=3,
entities_appeared=[{"id": "xiaoyan", "mentions": ["萧炎"], "confidence": 0.9}],
entities_new=[],
state_changes=[
{"entity_id": "xiaoyan", "field": "hp", "old": 5, "new": 0, "reason": "受伤"}
],
relationships_new=[
{"from_entity": "xiaoyan", "to_entity": "yaolao", "type": "师徒", "description": "收徒"}
],
)
assert stats["entities_updated"] >= 1
assert stats["state_changes"] == 1
updated = manager.get_entity("xiaoyan")
assert updated["current_json"]["hp"] == 0
rels = manager.get_entity_relationships("yaolao", direction="to")
assert rels
stats_summary = manager.get_stats()
assert "entities" in stats_summary
exported = manager.export_to_entities_v3_format()
assert exported["角色"]["xiaoyan"]["canonical_name"] == "萧炎"
def test_sql_state_manager_process_chapter_skips_and_existing(temp_project):
manager = SQLStateManager(temp_project)
manager.upsert_entity(EntityData(id="xiaoyan", type="角色", name="萧炎"))
stats = manager.process_chapter_entities(
chapter=1,
entities_appeared=[{"mentions": ["无ID"]}, {"id": "xiaoyan", "mentions": ["萧炎"]}],
entities_new=[{"name": "无ID"}, {"suggested_id": "xiaoyan", "name": "萧炎"}],
state_changes=[{"field": "realm"}, {"entity_id": "xiaoyan", "field": "hp", "old": 1, "new": 1}],
relationships_new=[{"from": "xiaoyan", "to": ""}],
)
assert stats["entities_updated"] >= 1
assert stats["relationships"] == 0
def test_sql_state_manager_export_protagonist_and_cli(temp_project, monkeypatch, capsys):
manager = SQLStateManager(temp_project)
def run_cli(args):
monkeypatch.setattr(sys, "argv", args)
sql_state_manager_module.main()
return json.loads(capsys.readouterr().out or "{}")
out = run_cli(["sql_state_manager", "--project-root", str(temp_project.project_root), "get-protagonist"])
assert out.get("status") == "error"
manager.upsert_entity(
EntityData(id="xiaoyan", type="角色", name="萧炎", is_protagonist=True)
)
exported = manager.export_to_entities_v3_format()
assert exported["角色"]["xiaoyan"]["is_protagonist"] is True
out = run_cli(["sql_state_manager", "--project-root", str(temp_project.project_root), "get-protagonist"])
assert out["status"] == "success"
assert out["data"].get("canonical_name") == "萧炎"
out = run_cli(["sql_state_manager", "--project-root", str(temp_project.project_root), "stats"])
assert out["status"] == "success"
assert "entities" in out.get("data", {})
out = run_cli(["sql_state_manager", "--project-root", str(temp_project.project_root), "get-core-entities"])
assert out["status"] == "success"
out = run_cli(["sql_state_manager", "--project-root", str(temp_project.project_root), "export-entities-v3"])
assert out["status"] == "success"
assert "角色" in out.get("data", {})
out = run_cli(["sql_state_manager", "--project-root", str(temp_project.project_root), "export-alias-index"])
assert out["status"] == "success"
assert isinstance(out.get("data", {}), dict)
payload = json.dumps({"entities_appeared": [], "entities_new": [], "state_changes": [], "relationships_new": []})
out = run_cli([
"sql_state_manager",
"--project-root",
str(temp_project.project_root),
"process-chapter",
"--chapter",
"2",
"--data",
payload,
])
assert out["status"] == "success"
@@ -0,0 +1,568 @@
#!/usr/bin/env python3
# -*- coding: utf-8 -*-
"""
StateManager extra tests
"""
import json
import sys
import tempfile
from pathlib import Path
import pytest
from data_modules.state_manager import StateManager, EntityState
from data_modules.index_manager import IndexManager, EntityMeta
@pytest.fixture
def temp_project(tmp_path):
from data_modules.config import DataModulesConfig
cfg = DataModulesConfig.from_project_root(tmp_path)
cfg.ensure_dirs()
return cfg
def test_ensure_state_schema_and_progress(temp_project):
# relationships as list should be migrated to structured_relationships
state = {
"relationships": [
{"from_entity": "a", "to_entity": "b", "type": "师徒", "chapter": 1}
],
"progress": {"current_chapter": "2", "total_words": "10"},
}
temp_project.state_file.write_text(json.dumps(state, ensure_ascii=False), encoding="utf-8")
manager = StateManager(temp_project, enable_sqlite_sync=False)
assert isinstance(manager._state.get("relationships"), dict)
assert isinstance(manager._state.get("structured_relationships"), list)
assert int(manager.get_current_chapter()) == 2
manager.update_progress(3)
assert manager.get_current_chapter() == 3
def test_add_update_entities_and_alias(temp_project):
manager = StateManager(temp_project, enable_sqlite_sync=False)
entity = EntityState(id="xiaoyan", name="萧炎", type="角色", tier="核心", aliases=["炎帝"])
assert manager.add_entity(entity) is True
assert manager.add_entity(entity) is False
manager.update_entity("xiaoyan", {"current": {"realm": "斗师"}})
updated = manager.get_entity("xiaoyan")
assert updated["current"]["realm"] == "斗师"
assert manager.get_entity_type("xiaoyan") == "角色"
assert manager.get_entity_type("missing") is None
assert "xiaoyan" in manager.get_all_entities()
assert "xiaoyan" in manager.get_entities_by_type("角色")
assert "xiaoyan" in manager.get_entities_by_tier("核心")
# unknown type update
assert manager.update_entity("missing", {"current": {"realm": "斗者"}}, "角色") is False
def test_update_entity_appearance_and_relationships(temp_project):
manager = StateManager(temp_project, enable_sqlite_sync=False)
manager.add_entity(EntityState(id="xiaoyan", name="萧炎", type="角色"))
manager.update_entity_appearance("xiaoyan", 5, "角色")
entity = manager.get_entity("xiaoyan")
assert entity.get("first_appearance") == 5
assert entity.get("last_appearance") == 5
# unknown entity should no-op
manager.update_entity_appearance("missing", 3, "角色")
manager.add_relationship("xiaoyan", "yaolao", "师徒", "收徒", 1)
rels = manager.get_relationships("xiaoyan")
assert len(rels) == 1
def test_disambiguation_and_save_state(temp_project):
manager = StateManager(temp_project, enable_sqlite_sync=False)
warnings = manager._record_disambiguation(
1,
[
{
"mention": "宗主",
"candidates": ["zongzhu", "lintian"],
"suggested": "zongzhu",
"confidence": 0.4,
},
{
"mention": "萧炎",
"candidates": [{"type": "角色", "id": "xiaoyan"}],
"suggested": "xiaoyan",
"confidence": 0.6,
},
],
)
assert any("需人工确认" in w for w in warnings)
assert any("消歧警告" in w for w in warnings)
manager.save_state()
assert temp_project.state_file.exists()
def test_save_state_no_pending(temp_project):
manager = StateManager(temp_project, enable_sqlite_sync=False)
manager.save_state()
assert not temp_project.state_file.exists()
def test_save_state_with_sqlite_sync_and_protagonist(temp_project):
manager = StateManager(temp_project)
manager.add_entity(EntityState(id="xiaoyan", name="萧炎", type="角色", tier="核心"))
manager.update_entity("xiaoyan", {"current": {"realm": "斗师", "location": "天云宗"}})
manager.update_progress(10, words=500)
manager.save_state()
state = json.loads(temp_project.state_file.read_text(encoding="utf-8"))
assert state.get("_migrated_to_sqlite") is True
assert state.get("progress", {}).get("current_chapter") == 10
# 标记为主角并同步
idx = IndexManager(temp_project)
idx.upsert_entity(
EntityMeta(
id="xiaoyan",
type="角色",
canonical_name="萧炎",
tier="核心",
current={"realm": "斗王", "location": "天云宗"},
first_appearance=1,
last_appearance=10,
is_protagonist=True,
),
update_metadata=True,
)
manager.sync_protagonist_from_entity()
assert manager._state.get("protagonist_state", {}).get("power", {}).get("realm") == "斗王"
manager._state["protagonist_state"] = {
"power": {"realm": "斗皇", "layer": 2},
"location": {"current": "中州"},
}
manager._state.setdefault("entities_v3", {"角色": {}})
manager._state["entities_v3"]["角色"]["xiaoyan"] = {
"canonical_name": "萧炎",
"tier": "核心",
"desc": "",
"current": {"realm": "斗王", "location": "天云宗"},
"first_appearance": 1,
"last_appearance": 10,
"history": [],
}
manager.sync_protagonist_to_entity("xiaoyan")
manager.save_state()
updated = idx.get_entity("xiaoyan")
assert updated["current_json"]["realm"] == "斗皇"
# export context
exported = manager.export_for_context()
assert exported.get("alias_index") == {}
def test_process_chapter_result_and_sqlite_sync(temp_project):
manager = StateManager(temp_project)
manager.add_entity(EntityState(id="xiaoyan", name="萧炎", type="角色", tier="核心"))
result = {
"entities_appeared": [
{"id": "xiaoyan", "type": "角色", "mentions": ["萧炎"], "confidence": 0.9}
],
"entities_new": [
{
"suggested_id": "yaolao",
"name": "药老",
"type": "角色",
"tier": "重要",
"mentions": ["药老"],
"aliases": ["药老先生"],
}
],
"state_changes": [
{"entity_id": "xiaoyan", "field": "realm", "old": "斗者", "new": "斗师", "reason": "突破"}
],
"relationships_new": [
{"from": "xiaoyan", "to": "yaolao", "type": "师徒", "description": "收徒"}
],
"uncertain": [
{"mention": "宗主", "candidates": ["zongzhu", "lintian"], "suggested": "zongzhu", "confidence": 0.2},
{
"mention": "萧炎",
"candidates": [{"type": "角色", "id": "xiaoyan"}],
"suggested": "xiaoyan",
"confidence": 0.8,
"adopted": True,
},
],
"chapter_meta": {"hook": "test", "end": "ok"},
}
warnings = manager.process_chapter_result(12, result)
assert any("需人工确认" in w for w in warnings)
assert any("消歧警告" in w for w in warnings)
manager.save_state()
idx = IndexManager(temp_project)
assert idx.get_entity("yaolao") is not None
assert idx.get_relationship_between("xiaoyan", "yaolao")
assert idx.get_entity_state_changes("xiaoyan")
by_type = manager.get_entities_by_type("角色")
by_tier = manager.get_entities_by_tier("核心")
assert "xiaoyan" in by_type
assert "xiaoyan" in by_tier
def test_export_context_and_protagonist_alias(temp_project):
manager = StateManager(temp_project, enable_sqlite_sync=False)
manager.add_entity(EntityState(id="xiaoyan", name="萧炎", type="角色", tier="核心"))
manager._state["disambiguation_warnings"] = [{"chapter": 1, "mention": "萧炎"}]
manager._state["disambiguation_pending"] = [{"chapter": 2, "mention": "宗主"}]
exported = manager.export_for_context()
assert "xiaoyan" in exported.get("entities", {})
assert exported["disambiguation"]["warnings"]
assert exported["disambiguation"]["pending"]
manager_sql = StateManager(temp_project)
idx = IndexManager(temp_project)
idx.upsert_entity(
EntityMeta(
id="xiaoyan",
type="角色",
canonical_name="萧炎",
tier="核心",
current={},
first_appearance=1,
last_appearance=1,
is_protagonist=False,
),
update_metadata=True,
)
idx.register_alias("小炎子", "xiaoyan", "角色")
manager_sql._state["protagonist_state"] = {"name": "小炎子"}
assert manager_sql.get_protagonist_entity_id() == "xiaoyan"
idx.upsert_entity(
EntityMeta(
id="xiaoyan",
type="角色",
canonical_name="萧炎",
tier="核心",
current={},
first_appearance=1,
last_appearance=1,
is_protagonist=True,
),
update_metadata=True,
)
assert manager_sql.get_protagonist_entity_id() == "xiaoyan"
def test_sqlite_metadata_update_and_alias_sync(temp_project):
manager = StateManager(temp_project)
idx = IndexManager(temp_project)
idx.upsert_entity(
EntityMeta(
id="xiaoyan",
type="角色",
canonical_name="萧炎",
tier="核心",
current={"realm": "斗者"},
first_appearance=1,
last_appearance=1,
is_protagonist=False,
)
)
manager._state.setdefault("entities_v3", {"角色": {}})
manager._state["entities_v3"]["角色"]["xiaoyan"] = {
"canonical_name": "萧炎",
"tier": "核心",
"desc": "",
"current": {"realm": "斗者"},
"first_appearance": 1,
"last_appearance": 1,
"history": [],
}
manager.update_entity(
"xiaoyan",
{"canonical_name": "萧炎·新", "tier": "重要", "current": {"realm": "斗王"}},
"角色",
)
manager.update_entity("xiaoyan", {"location": "中州"}, "角色")
manager.update_entity_appearance("xiaoyan", 2, "角色")
manager._pending_alias_entries["小炎子"] = [{"type": "角色", "id": "xiaoyan"}]
manager.save_state()
updated = idx.get_entity("xiaoyan")
assert updated["canonical_name"] == "萧炎·新"
assert updated["current_json"]["realm"] == "斗王"
assert updated["current_json"]["location"] == "中州"
assert updated["last_appearance"] == 2
aliases = idx.get_entity_aliases("xiaoyan")
assert "萧炎·新" in aliases
assert "小炎子" in aliases
def test_ensure_state_schema_invalid_inputs(temp_project):
manager = StateManager(temp_project, enable_sqlite_sync=False)
schema = manager._ensure_state_schema("bad")
assert isinstance(schema, dict)
schema2 = manager._ensure_state_schema({
"progress": "bad",
"relationships": "bad",
"disambiguation_warnings": "bad",
"disambiguation_pending": "bad",
})
assert isinstance(schema2["progress"], dict)
assert isinstance(schema2["relationships"], dict)
assert isinstance(schema2["disambiguation_warnings"], list)
assert isinstance(schema2["disambiguation_pending"], list)
def test_save_state_preserves_sqlite_pending_on_sync_failure(temp_project):
manager = StateManager(temp_project)
manager.add_entity(EntityState(id="e1", name="测试角色", type="角色", first_appearance=1, last_appearance=1))
manager.update_entity("e1", {"current": {"realm": "炼气"}}, "角色")
class _BrokenSQLManager:
def process_chapter_entities(self, **kwargs):
raise RuntimeError("boom")
manager._sql_state_manager = _BrokenSQLManager()
manager._pending_sqlite_data["chapter"] = 1
manager.save_state()
state = json.loads(temp_project.state_file.read_text(encoding="utf-8"))
assert state.get("_migrated_to_sqlite") is True
# SQLite 同步失败后,SQLite 相关 pending 不应被清空,便于后续重试
assert manager._pending_entity_patches
assert manager._pending_sqlite_data.get("chapter") == 1
def test_save_state_progress_and_disambiguation_merge(temp_project):
state = {
"progress": {"current_chapter": "bad", "total_words": "bad"},
"disambiguation_warnings": "bad",
"disambiguation_pending": "bad",
}
temp_project.state_file.write_text(json.dumps(state, ensure_ascii=False), encoding="utf-8")
manager = StateManager(temp_project, enable_sqlite_sync=False)
manager.config.max_disambiguation_warnings = 1
manager.config.max_disambiguation_pending = 1
manager._pending_progress_chapter = 5
manager._pending_progress_words_delta = 10
manager._pending_disambiguation_warnings = [
{"chapter": 1, "mention": "a", "chosen_id": "x", "confidence": 0.5},
{"chapter": 1, "mention": "a", "chosen_id": "x", "confidence": 0.5},
"bad",
]
manager._pending_disambiguation_pending = [
{"chapter": 2, "mention": "b", "suggested_id": "y", "confidence": 0.4},
{"chapter": 2, "mention": "b", "suggested_id": "y", "confidence": 0.4},
"bad",
]
manager.save_state()
saved = json.loads(temp_project.state_file.read_text(encoding="utf-8"))
assert saved["progress"]["current_chapter"] == 5
assert saved["progress"]["total_words"] == 10
assert len(saved["disambiguation_warnings"]) == 1
assert len(saved["disambiguation_pending"]) == 1
def test_sync_to_sqlite_exceptions_and_no_sql_manager(temp_project, monkeypatch):
manager = StateManager(temp_project)
manager._pending_progress_chapter = 1
manager._pending_sqlite_data["chapter"] = 1
manager._pending_alias_entries["alias"] = [{"type": "角色", "id": "xiaoyan"}]
def boom(*args, **kwargs):
raise RuntimeError("boom")
monkeypatch.setattr(manager._sql_state_manager, "process_chapter_entities", boom)
monkeypatch.setattr(manager._sql_state_manager, "register_alias", boom)
manager.save_state()
manager_no_sql = StateManager(temp_project, enable_sqlite_sync=False)
manager_no_sql._sync_pending_patches_to_sqlite()
def test_entity_fallbacks_and_updates(temp_project):
manager = StateManager(temp_project, enable_sqlite_sync=False)
manager.add_entity(EntityState(id="hero", name="主角", type="未知", tier="核心"))
manager.add_entity(EntityState(id="place", name="乌坦城", type="地点", tier="重要"))
assert manager.get_entity("hero", "角色")["canonical_name"] == "主角"
assert manager.get_entity("place")["canonical_name"] == "乌坦城"
assert manager.get_entity_type("place") == "地点"
assert "hero" in manager.get_entities_by_type("角色")
assert "hero" in manager.get_entities_by_tier("核心")
assert "hero" in manager.get_all_entities()
assert manager.update_entity("missing", {"current": {"a": 1}}) is False
manager.update_entity("hero", {"attributes": {"hp": 1}}, "角色")
manager._state["entities_v3"]["角色"]["hero"].pop("current", None)
manager.update_entity("hero", {"current": {"mp": 2}}, "角色")
manager.update_entity("hero", {"tier": "重要"}, "角色")
manager._state["entities_v3"] = "bad"
manager.update_entity_appearance("hero", 1, "角色")
manager._state["entities_v3"]["角色"]["hero"] = {"first_appearance": 0, "last_appearance": 0}
manager.update_entity_appearance("hero", 1, "角色")
manager.update_entity_appearance("hero", 2, "角色")
def test_register_alias_internal_and_get_all_entities_sqlite(temp_project):
manager = StateManager(temp_project)
manager._register_alias_internal("xiaoyan", "角色", "")
manager._register_alias_internal("xiaoyan", "角色", "萧炎")
idx = IndexManager(temp_project)
idx.upsert_entity(
EntityMeta(
id="xiaoyan",
type="角色",
canonical_name="萧炎",
tier="核心",
current={},
first_appearance=1,
last_appearance=1,
is_protagonist=False,
)
)
all_entities = manager.get_all_entities()
assert "xiaoyan" in all_entities
def test_record_disambiguation_and_process_chapter_existing(temp_project):
manager = StateManager(temp_project, enable_sqlite_sync=False)
warnings = manager._record_disambiguation(
1,
[
"bad",
{"mention": "", "confidence": 0.1},
{"mention": "宗主", "confidence": "bad", "adopted": "zongzhu"},
],
)
assert warnings
manager.add_entity(EntityState(id="xiaoyan", name="萧炎", type="角色"))
warnings = manager.process_chapter_result(2, {"entities_new": [{"id": "xiaoyan", "name": "萧炎"}]})
assert any("实体已存在" in w for w in warnings)
def test_sync_protagonist_from_string_and_empty_updates(temp_project):
manager = StateManager(temp_project, enable_sqlite_sync=False)
manager._state.setdefault("entities_v3", {"角色": {}})
manager._state["entities_v3"]["角色"]["bad"] = {
"current": None,
"current_json": "not-json",
}
manager._state["entities_v3"]["角色"]["hero"] = {
"current": None,
"current_json": json.dumps({"realm": "斗师", "layer": 2, "location": "乌坦城", "last_chapter": 3}),
}
manager.sync_protagonist_from_entity("bad")
manager.sync_protagonist_from_entity("hero")
assert manager._state["protagonist_state"]["power"]["realm"] == "斗师"
manager._state["protagonist_state"] = {}
manager.sync_protagonist_to_entity()
def test_state_manager_cli_commands(temp_project, monkeypatch, capsys):
idx = IndexManager(temp_project)
idx.upsert_entity(
EntityMeta(
id="xiaoyan",
type="角色",
canonical_name="萧炎",
tier="核心",
current={},
first_appearance=1,
last_appearance=1,
is_protagonist=False,
)
)
def run_cli(args):
monkeypatch.setattr(sys, "argv", args)
from data_modules import state_manager as sm
sm.main()
out = capsys.readouterr().out
return json.loads(out)
out = run_cli(["state_manager", "--project-root", str(temp_project.project_root), "get-progress"])
assert out["status"] == "success"
assert "current_chapter" in out.get("data", {})
out = run_cli(["state_manager", "--project-root", str(temp_project.project_root), "get-entity", "--id", "missing"])
assert out["status"] == "error"
out = run_cli(["state_manager", "--project-root", str(temp_project.project_root), "get-entity", "--id", "xiaoyan"])
assert out["status"] == "success"
assert out["data"].get("id") == "xiaoyan"
out = run_cli(["state_manager", "--project-root", str(temp_project.project_root), "list-entities", "--type", "角色"])
assert out["status"] == "success"
assert any(e.get("id") == "xiaoyan" for e in out.get("data", []))
out = run_cli(["state_manager", "--project-root", str(temp_project.project_root), "list-entities", "--tier", "核心"])
assert out["status"] == "success"
assert any(e.get("id") == "xiaoyan" for e in out.get("data", []))
payload = json.dumps({"entities_appeared": [], "entities_new": [], "state_changes": [], "relationships_new": []})
out = run_cli([
"state_manager",
"--project-root",
str(temp_project.project_root),
"process-chapter",
"--chapter",
"1",
"--data",
payload,
])
assert out["status"] == "success"
def test_save_state_timeout(monkeypatch, temp_project):
import filelock
from data_modules import state_manager as sm
manager = StateManager(temp_project, enable_sqlite_sync=False)
manager.update_progress(1)
class FakeLock:
def __init__(self, *args, **kwargs):
pass
def __enter__(self):
raise filelock.Timeout("timeout")
def __exit__(self, exc_type, exc, tb):
return False
monkeypatch.setattr(sm.filelock, "FileLock", FakeLock)
with pytest.raises(RuntimeError):
manager.save_state()
@@ -0,0 +1,107 @@
#!/usr/bin/env python3
# -*- coding: utf-8 -*-
from data_modules.state_validator import (
FORESHADOWING_STATUS_PENDING,
FORESHADOWING_STATUS_RESOLVED,
FORESHADOWING_TIER_CORE,
FORESHADOWING_TIER_DECOR,
FORESHADOWING_TIER_SUB,
count_patterns,
get_chapter_meta_entry,
is_resolved_foreshadowing_status,
normalize_chapter_meta,
normalize_foreshadowing_item,
normalize_foreshadowing_status,
normalize_foreshadowing_tier,
normalize_state_runtime_sections,
resolve_chapter_field,
split_patterns,
to_positive_int,
)
def test_to_positive_int_and_resolve_chapter_field():
assert to_positive_int(12) == 12
assert to_positive_int("ch-18") == 18
assert to_positive_int(0) is None
assert to_positive_int("no number") is None
item = {"added_chapter": "第15章", "target": "200"}
assert resolve_chapter_field(item, ["planted_chapter", "added_chapter"]) == 15
assert resolve_chapter_field(item, ["target_chapter", "target"]) == 200
def test_status_and_tier_normalization():
assert normalize_foreshadowing_status("pending") == FORESHADOWING_STATUS_PENDING
assert normalize_foreshadowing_status("resolved") == FORESHADOWING_STATUS_RESOLVED
assert normalize_foreshadowing_status("") == FORESHADOWING_STATUS_PENDING
assert is_resolved_foreshadowing_status("已回收") is True
assert is_resolved_foreshadowing_status("active") is False
assert normalize_foreshadowing_tier("core") == FORESHADOWING_TIER_CORE
assert normalize_foreshadowing_tier("decoration") == FORESHADOWING_TIER_DECOR
assert normalize_foreshadowing_tier("unknown") == FORESHADOWING_TIER_SUB
def test_pattern_split_and_count():
assert split_patterns(["A", " A ", "B", ""]) == ["A", "B"]
assert split_patterns("A, B / C|A") == ["A", "B", "C"]
assert count_patterns("A,B,C") == 3
assert count_patterns(123) is None
def test_normalize_foreshadowing_item_and_chapter_meta_entry():
item = {
"content": " 遗迹钥匙 ",
"status": "pending",
"tier": "main",
"added_chapter": "第30章",
"target": "120",
}
normalized_item = normalize_foreshadowing_item(item)
assert normalized_item["content"] == "遗迹钥匙"
assert normalized_item["status"] == FORESHADOWING_STATUS_PENDING
assert normalized_item["tier"] == FORESHADOWING_TIER_CORE
assert normalized_item["planted_chapter"] == 30
assert normalized_item["target_chapter"] == 120
state = {
"chapter_meta": {
"0003": {"coolpoint_pattern": "反杀, 掉马"},
"7": {"patterns": ["翻车", "反杀"]},
}
}
meta3 = get_chapter_meta_entry(state, 3)
assert meta3["coolpoint_patterns"] == ["反杀", "掉马"]
meta7 = get_chapter_meta_entry(state, 7)
assert meta7["coolpoint_patterns"] == ["翻车", "反杀"]
def test_normalize_state_runtime_sections():
state = {
"plot_threads": {
"foreshadowing": [
{"content": "伏笔A", "status": "active", "tier": "decor", "chapter": 11, "target": 99},
"invalid",
]
},
"chapter_meta": {
1: {"cool_point_pattern": "打脸|翻车"},
"bad": "invalid",
},
}
normalized = normalize_state_runtime_sections(state)
assert len(normalized["plot_threads"]["foreshadowing"]) == 1
first = normalized["plot_threads"]["foreshadowing"][0]
assert first["status"] == FORESHADOWING_STATUS_PENDING
assert first["tier"] == FORESHADOWING_TIER_DECOR
assert first["planted_chapter"] == 11
assert first["target_chapter"] == 99
chapter_meta = normalize_chapter_meta(normalized["chapter_meta"])
assert "1" in chapter_meta
assert chapter_meta["1"]["coolpoint_patterns"] == ["打脸", "翻车"]
@@ -0,0 +1,235 @@
#!/usr/bin/env python3
# -*- coding: utf-8 -*-
import json
import tempfile
from data_modules.config import DataModulesConfig
from data_modules.index_manager import (
IndexManager,
ChapterReadingPowerMeta,
EntityMeta,
RelationshipMeta,
RelationshipEventMeta,
)
from status_reporter import StatusReporter
def _write_state(project_root, state: dict):
noma_dir = project_root / ".noma"
noma_dir.mkdir(parents=True, exist_ok=True)
(noma_dir / "state.json").write_text(
json.dumps(state, ensure_ascii=False, indent=2),
encoding="utf-8",
)
def test_foreshadowing_analysis_uses_real_chapters_and_handles_missing_data():
with tempfile.TemporaryDirectory() as tmpdir:
project_root = DataModulesConfig.from_project_root(tmpdir).project_root
state = {
"progress": {"current_chapter": 120, "total_words": 360000},
"plot_threads": {
"foreshadowing": [
{
"content": "林家宝库铭文的秘密",
"status": "未回收",
"tier": "核心",
"planted_chapter": 20,
"target_chapter": 100,
},
{
"content": "神秘玉佩来历",
"status": "待回收",
"tier": "支线",
"added_chapter": 50,
"target": 150,
},
{
"content": "旧日誓言",
"status": "未回收",
"tier": "装饰",
},
{
"content": "已完成伏笔",
"status": "已回收",
"planted_chapter": 10,
"target_chapter": 20,
},
]
},
}
_write_state(project_root, state)
reporter = StatusReporter(str(project_root))
assert reporter.load_state() is True
foreshadowing = reporter.analyze_foreshadowing()
assert len(foreshadowing) == 3
records = {item["content"]: item for item in foreshadowing}
assert records["林家宝库铭文的秘密"]["planted_chapter"] == 20
assert records["林家宝库铭文的秘密"]["elapsed"] == 100
assert records["林家宝库铭文的秘密"]["status"] == "🔴 已超期"
assert records["神秘玉佩来历"]["planted_chapter"] == 50
assert records["神秘玉佩来历"]["target_chapter"] == 150
assert records["神秘玉佩来历"]["status"] in {"🟡 轻度超时", "🟢 正常"}
assert records["旧日誓言"]["planted_chapter"] is None
assert records["旧日誓言"]["status"] == "⚪ 数据不足"
urgency = reporter.analyze_foreshadowing_urgency()
urgency_by_content = {item["content"]: item for item in urgency}
assert urgency_by_content["林家宝库铭文的秘密"]["urgency"] is not None
assert urgency_by_content["林家宝库铭文的秘密"]["status"] == "🔴 已超期"
assert urgency_by_content["旧日誓言"]["urgency"] is None
assert urgency_by_content["旧日誓言"]["status"] == "⚪ 数据不足"
def test_pacing_analysis_prefers_real_coolpoint_metadata_over_estimation():
with tempfile.TemporaryDirectory() as tmpdir:
config = DataModulesConfig.from_project_root(tmpdir)
config.ensure_dirs()
project_root = config.project_root
state = {
"progress": {"current_chapter": 3, "total_words": 12000},
"chapter_meta": {
"0003": {
"hook": "下章有变",
"coolpoint_patterns": ["身份掉马", "反派翻车"],
}
},
}
_write_state(project_root, state)
idx = IndexManager(config)
idx.save_chapter_reading_power(
ChapterReadingPowerMeta(
chapter=1,
hook_type="渴望钩",
hook_strength="strong",
coolpoint_patterns=["打脸权威", "身份掉马"],
)
)
idx.save_chapter_reading_power(
ChapterReadingPowerMeta(
chapter=2,
hook_type="悬念钩",
hook_strength="medium",
coolpoint_patterns=["身份掉马"],
)
)
reporter = StatusReporter(str(project_root))
assert reporter.load_state() is True
reporter.chapters_data = [
{"chapter": 1, "word_count": 4000, "cool_point": "", "dominant": "", "characters": []},
{"chapter": 2, "word_count": 3000, "cool_point": "", "dominant": "", "characters": []},
{"chapter": 3, "word_count": 5000, "cool_point": "", "dominant": "", "characters": []},
]
segments = reporter.analyze_pacing()
assert len(segments) == 1
seg = segments[0]
assert seg["cool_points"] == 5
assert round(seg["words_per_point"], 2) == 2400.00
assert seg["missing_chapters"] == 0
assert seg["dominant_source"] == "chapter_reading_power"
def test_pacing_analysis_marks_missing_data_instead_of_assuming_one_point_per_chapter():
with tempfile.TemporaryDirectory() as tmpdir:
config = DataModulesConfig.from_project_root(tmpdir)
config.ensure_dirs()
project_root = config.project_root
state = {
"progress": {"current_chapter": 1, "total_words": 2000},
"chapter_meta": {},
}
_write_state(project_root, state)
reporter = StatusReporter(str(project_root))
assert reporter.load_state() is True
reporter.chapters_data = [
{"chapter": 1, "word_count": 2000, "cool_point": "", "dominant": "", "characters": []}
]
seg = reporter.analyze_pacing()[0]
assert seg["cool_points"] == 0
assert seg["words_per_point"] is None
assert seg["rating"] == "数据不足"
assert seg["missing_chapters"] == 1
def test_relationship_graph_prefers_index_db_data():
with tempfile.TemporaryDirectory() as tmpdir:
config = DataModulesConfig.from_project_root(tmpdir)
config.ensure_dirs()
project_root = config.project_root
state = {
"progress": {"current_chapter": 12, "total_words": 24000},
"protagonist_state": {"name": "萧炎"},
"relationships": {"allies": [{"name": "旧盟友", "relation": "友好"}], "enemies": []},
}
_write_state(project_root, state)
idx = IndexManager(config)
idx.upsert_entity(
EntityMeta(
id="xiaoyan",
type="角色",
canonical_name="萧炎",
tier="核心",
current={},
first_appearance=1,
last_appearance=12,
is_protagonist=True,
)
)
idx.upsert_entity(
EntityMeta(
id="yaolao",
type="角色",
canonical_name="药老",
tier="重要",
current={},
first_appearance=1,
last_appearance=12,
)
)
idx.upsert_relationship(
RelationshipMeta(
from_entity="xiaoyan",
to_entity="yaolao",
type="师徒",
description="师徒关系",
chapter=10,
)
)
idx.record_relationship_event(
RelationshipEventMeta(
from_entity="xiaoyan",
to_entity="yaolao",
type="师徒",
chapter=10,
action="create",
polarity=1,
strength=0.9,
description="拜师",
evidence="萧炎拜药老为师",
)
)
reporter = StatusReporter(str(project_root))
assert reporter.load_state() is True
graph = reporter.generate_relationship_graph()
assert "mermaid" in graph
assert "药老" in graph
assert "师徒" in graph
@@ -0,0 +1,91 @@
#!/usr/bin/env python3
# -*- coding: utf-8 -*-
"""
StyleSampler extra tests + CLI
"""
import sys
import json
import pytest
import data_modules.style_sampler as sampler_module
from data_modules.style_sampler import StyleSampler, StyleSample, SceneType
from data_modules.config import DataModulesConfig
@pytest.fixture
def temp_project(tmp_path):
cfg = DataModulesConfig.from_project_root(tmp_path)
cfg.ensure_dirs()
return cfg
def test_style_sampler_more(temp_project):
sampler = StyleSampler(temp_project)
sample = StyleSample(
id="ch1_s1",
chapter=1,
scene_type=SceneType.BATTLE.value,
content="战斗描写很精彩",
score=0.9,
tags=["战斗"],
)
assert sampler.add_sample(sample) is True
assert sampler.add_sample(sample) is False
best = sampler.get_best_samples(limit=5)
assert len(best) == 1
stats = sampler.get_stats()
assert stats["total"] == 1
# scene type inference
assert sampler._infer_scene_types("一场战斗") == [SceneType.BATTLE.value]
assert sampler._infer_scene_types("对话和谈话") == [SceneType.DIALOGUE.value]
assert sampler._infer_scene_types("心理情感描写") == [SceneType.EMOTION.value]
# classify and tags
scene_type = sampler._classify_scene_type({"summary": "紧张", "content": ""})
assert scene_type == SceneType.TENSION.value
tags = sampler._extract_tags("战斗 修炼 对话 描写")
assert "战斗" in tags
def test_style_sampler_cli(temp_project, monkeypatch, capsys):
root = str(temp_project.project_root)
def run_cli(args):
monkeypatch.setattr(sys, "argv", ["style_sampler"] + args)
sampler_module.main()
run_cli(["--project-root", root, "stats"])
run_cli(["--project-root", root, "list", "--limit", "5"])
run_cli(
[
"--project-root",
root,
"extract",
"--chapter",
"1",
"--score",
"90",
"--scenes",
json.dumps(
[
{
"index": 1,
"summary": "战斗场景",
"content": "战斗" + "a" * 300,
}
],
ensure_ascii=False,
),
]
)
run_cli(["--project-root", root, "list", "--type", "战斗", "--limit", "5"])
run_cli(["--project-root", root, "select", "--outline", "本章有一场战斗", "--max", "2"])
capsys.readouterr()
@@ -0,0 +1,45 @@
#!/usr/bin/env python3
# -*- coding: utf-8 -*-
import json
import sys
def test_update_state_cli_add_review_writes_checkpoint(tmp_path, monkeypatch):
import update_state as update_state_module
noma_dir = tmp_path / ".noma"
noma_dir.mkdir(parents=True, exist_ok=True)
state = {
"project_info": {},
"progress": {"current_chapter": 1, "total_words": 0},
"protagonist_state": {
"power": {"realm": "炼气", "layer": 1, "bottleneck": None},
"location": "村口",
},
"relationships": {},
"world_settings": {},
"plot_threads": {},
"review_checkpoints": [],
}
state_file = noma_dir / "state.json"
state_file.write_text(json.dumps(state, ensure_ascii=False), encoding="utf-8")
# 避免在测试里创建备份目录/修改权限等非核心行为
monkeypatch.setattr(update_state_module.StateUpdater, "backup", lambda self: True)
report_file = "review/report_1_2.md"
monkeypatch.setattr(
sys,
"argv",
["update_state", "--project-root", str(tmp_path), "--add-review", "1-2", report_file],
)
update_state_module.main()
updated = json.loads(state_file.read_text(encoding="utf-8"))
checkpoints = updated.get("review_checkpoints")
assert isinstance(checkpoints, list)
assert checkpoints[-1]["chapters"] == "1-2"
assert checkpoints[-1]["report"] == report_file
@@ -0,0 +1,168 @@
#!/usr/bin/env python3
# -*- coding: utf-8 -*-
import sys
from pathlib import Path
import pytest
def _ensure_scripts_on_path() -> None:
scripts_dir = Path(__file__).resolve().parents[2]
if str(scripts_dir) not in sys.path:
sys.path.insert(0, str(scripts_dir))
def _load_noma_module():
_ensure_scripts_on_path()
import data_modules.noma as noma_module
return noma_module
def test_init_does_not_resolve_existing_project_root(monkeypatch):
module = _load_noma_module()
called = {}
def _fake_run_script(script_name, argv):
called["script_name"] = script_name
called["argv"] = list(argv)
return 0
def _fail_resolve(_explicit_project_root=None):
raise AssertionError("init 子命令不应触发 project_root 解析")
monkeypatch.setenv("WEBNOVEL_PROJECT_ROOT", r"D:\invalid\root")
monkeypatch.setattr(module, "_run_script", _fake_run_script)
monkeypatch.setattr(module, "_resolve_root", _fail_resolve)
monkeypatch.setattr(sys, "argv", ["noma", "init", "proj-dir", "测试书", "修仙"])
with pytest.raises(SystemExit) as exc:
module.main()
assert int(exc.value.code or 0) == 0
assert called["script_name"] == "init_project.py"
assert called["argv"] == ["proj-dir", "测试书", "修仙"]
def test_extract_context_forwards_with_resolved_project_root(monkeypatch, tmp_path):
module = _load_noma_module()
book_root = (tmp_path / "book").resolve()
called = {}
def _fake_resolve(explicit_project_root=None):
return book_root
def _fake_run_script(script_name, argv):
called["script_name"] = script_name
called["argv"] = list(argv)
return 0
monkeypatch.setattr(module, "_resolve_root", _fake_resolve)
monkeypatch.setattr(module, "_run_script", _fake_run_script)
monkeypatch.setattr(
sys,
"argv",
[
"noma",
"--project-root",
str(tmp_path),
"extract-context",
"--chapter",
"12",
"--format",
"json",
],
)
with pytest.raises(SystemExit) as exc:
module.main()
assert int(exc.value.code or 0) == 0
assert called["script_name"] == "extract_chapter_context.py"
assert called["argv"] == [
"--project-root",
str(book_root),
"--chapter",
"12",
"--format",
"json",
]
def test_preflight_succeeds_for_valid_project_root(monkeypatch, tmp_path, capsys):
module = _load_noma_module()
project_root = tmp_path / "book"
(project_root / ".noma").mkdir(parents=True, exist_ok=True)
(project_root / ".noma" / "state.json").write_text("{}", encoding="utf-8")
monkeypatch.setattr(sys, "argv", ["noma", "--project-root", str(project_root), "preflight"])
with pytest.raises(SystemExit) as exc:
module.main()
captured = capsys.readouterr()
assert int(exc.value.code or 0) == 0
assert "OK project_root" in captured.out
assert str(project_root.resolve()) in captured.out
def test_preflight_fails_when_required_scripts_are_missing(monkeypatch, tmp_path, capsys):
module = _load_noma_module()
project_root = tmp_path / "book"
(project_root / ".noma").mkdir(parents=True, exist_ok=True)
(project_root / ".noma" / "state.json").write_text("{}", encoding="utf-8")
fake_scripts_dir = tmp_path / "fake-scripts"
fake_scripts_dir.mkdir(parents=True, exist_ok=True)
monkeypatch.setattr(module, "_scripts_dir", lambda: fake_scripts_dir)
monkeypatch.setattr(sys, "argv", ["noma", "--project-root", str(project_root), "preflight", "--format", "json"])
with pytest.raises(SystemExit) as exc:
module.main()
captured = capsys.readouterr()
assert int(exc.value.code or 0) == 1
assert '"ok": false' in captured.out
assert '"name": "entry_script"' in captured.out
def test_quality_trend_report_writes_to_book_root_when_input_is_workspace_root(tmp_path, monkeypatch):
_ensure_scripts_on_path()
import quality_trend_report as quality_trend_report_module
workspace_root = (tmp_path / "workspace").resolve()
book_root = (workspace_root / "凡人资本论").resolve()
(workspace_root / ".claude").mkdir(parents=True, exist_ok=True)
(workspace_root / ".claude" / ".noma-current-project").write_text(str(book_root), encoding="utf-8")
(book_root / ".noma").mkdir(parents=True, exist_ok=True)
(book_root / ".noma" / "state.json").write_text("{}", encoding="utf-8")
output_path = workspace_root / "report.md"
monkeypatch.setattr(
sys,
"argv",
[
"quality_trend_report",
"--project-root",
str(workspace_root),
"--limit",
"1",
"--output",
str(output_path),
],
)
quality_trend_report_module.main()
assert output_path.is_file()
assert (book_root / ".noma" / "index.db").is_file()
assert not (workspace_root / ".noma" / "index.db").exists()
@@ -0,0 +1,204 @@
#!/usr/bin/env python3
# -*- coding: utf-8 -*-
import json
import logging
import sys
from pathlib import Path
from types import SimpleNamespace
def _load_module():
scripts_dir = Path(__file__).resolve().parents[2]
if str(scripts_dir) not in sys.path:
sys.path.insert(0, str(scripts_dir))
import workflow_manager
return workflow_manager
def test_workflow_lifecycle_and_trace(tmp_path, monkeypatch):
module = _load_module()
monkeypatch.setattr(module, "find_project_root", lambda: tmp_path)
noma_dir = tmp_path / ".noma"
noma_dir.mkdir(parents=True, exist_ok=True)
module.start_task("noma-write", {"chapter_num": 7})
module.start_step("Step 1", "Context")
module.complete_step("Step 1", json.dumps({"state_json_modified": True}, ensure_ascii=False))
module.complete_task(json.dumps({"review_completed": True}, ensure_ascii=False))
state = module.load_state()
assert state["current_task"] is None
assert state["history"][-1]["status"] == module.TASK_STATUS_COMPLETED
assert state["last_stable_state"]["artifacts"]["review_completed"] is True
trace_path = module.get_call_trace_path()
assert trace_path.exists()
lines = trace_path.read_text(encoding="utf-8").strip().splitlines()
events = [json.loads(line)["event"] for line in lines if line.strip()]
assert "task_started" in events
assert "step_started" in events
assert "step_completed" in events
assert "task_completed" in events
def test_start_task_reentry_increments_retry(tmp_path, monkeypatch):
module = _load_module()
monkeypatch.setattr(module, "find_project_root", lambda: tmp_path)
noma_dir = tmp_path / ".noma"
noma_dir.mkdir(parents=True, exist_ok=True)
module.start_task("noma-write", {"chapter_num": 8})
module.start_task("noma-write", {"chapter_num": 8})
state = module.load_state()
task = state["current_task"]
assert task is not None
assert task["status"] == module.TASK_STATUS_RUNNING
assert int(task.get("retry_count", 0)) >= 1
def test_complete_step_rejects_mismatch_step_id(tmp_path, monkeypatch):
module = _load_module()
monkeypatch.setattr(module, "find_project_root", lambda: tmp_path)
noma_dir = tmp_path / ".noma"
noma_dir.mkdir(parents=True, exist_ok=True)
module.start_task("noma-write", {"chapter_num": 9})
module.start_step("Step 2A", "Draft")
module.complete_step("Step 2B")
state = module.load_state()
current_step = state["current_task"]["current_step"]
assert current_step is not None
assert current_step["id"] == "Step 2A"
assert current_step["status"] == module.STEP_STATUS_RUNNING
def test_workflow_step_owner_and_order_violation_trace(tmp_path, monkeypatch):
module = _load_module()
monkeypatch.setattr(module, "find_project_root", lambda: tmp_path)
noma_dir = tmp_path / ".noma"
noma_dir.mkdir(parents=True, exist_ok=True)
assert module.expected_step_owner("noma-write", "Step 1") == "context-agent"
assert module.expected_step_owner("noma-write", "Step 5") == "data-agent"
module.start_task("noma-write", {"chapter_num": 12})
module.start_step("Step 3", "Review")
trace_path = module.get_call_trace_path()
lines = [json.loads(line) for line in trace_path.read_text(encoding="utf-8").splitlines() if line.strip()]
events = [row.get("event") for row in lines]
assert "step_order_violation" in events
step_started = [row for row in lines if row.get("event") == "step_started"]
assert step_started
assert step_started[-1].get("payload", {}).get("expected_owner") == "review-agents"
def test_safe_append_call_trace_logs_failure(monkeypatch, caplog):
module = _load_module()
def _raise_trace_error(event, payload=None):
raise RuntimeError("trace failure")
monkeypatch.setattr(module, "append_call_trace", _raise_trace_error)
with caplog.at_level(logging.WARNING):
module.safe_append_call_trace("unit_test_event", {"ok": True})
message_text = "\n".join(record.getMessage() for record in caplog.records)
assert "failed to append call trace" in message_text
assert "unit_test_event" in message_text
def test_get_workflow_paths_support_zero_arg_find_project_root(tmp_path, monkeypatch):
module = _load_module()
monkeypatch.setattr(module, "_cli_project_root", None)
monkeypatch.setattr(module, "find_project_root", lambda: tmp_path)
assert module.get_workflow_state_path() == tmp_path / ".noma" / "workflow_state.json"
assert module.get_call_trace_path() == tmp_path / ".noma" / "observability" / "call_trace.jsonl"
def test_workflow_reentry_does_not_duplicate_history(tmp_path, monkeypatch):
module = _load_module()
monkeypatch.setattr(module, "find_project_root", lambda: tmp_path)
noma_dir = tmp_path / ".noma"
noma_dir.mkdir(parents=True, exist_ok=True)
module.start_task("noma-write", {"chapter_num": 20})
module.start_task("noma-write", {"chapter_num": 20})
module.start_task("noma-write", {"chapter_num": 20})
state = module.load_state()
assert isinstance(state.get("history"), list)
assert len(state.get("history")) == 0
task = state.get("current_task") or {}
assert int(task.get("retry_count", 0)) >= 2
def test_cleanup_artifacts_requires_confirm(tmp_path, monkeypatch):
module = _load_module()
monkeypatch.setattr(module, "find_project_root", lambda: tmp_path)
noma_dir = tmp_path / ".noma"
noma_dir.mkdir(parents=True, exist_ok=True)
draft_path = module.default_chapter_draft_path(tmp_path, 7)
draft_path.parent.mkdir(parents=True, exist_ok=True)
draft_path.write_text("draft", encoding="utf-8")
git_called = {"count": 0}
def _fake_run(*args, **kwargs):
git_called["count"] += 1
return SimpleNamespace(returncode=0, stderr="", stdout="")
monkeypatch.setattr(module.subprocess, "run", _fake_run)
preview = module.cleanup_artifacts(7, confirm=False)
assert draft_path.exists()
assert git_called["count"] == 0
assert any(item.startswith("[预览]") for item in preview)
def test_cleanup_artifacts_confirm_deletes_with_backup(tmp_path, monkeypatch):
module = _load_module()
monkeypatch.setattr(module, "find_project_root", lambda: tmp_path)
noma_dir = tmp_path / ".noma"
noma_dir.mkdir(parents=True, exist_ok=True)
draft_path = module.default_chapter_draft_path(tmp_path, 8)
draft_path.parent.mkdir(parents=True, exist_ok=True)
draft_path.write_text("draft", encoding="utf-8")
git_called = {"count": 0, "cmd": None}
def _fake_run(cmd, **kwargs):
git_called["count"] += 1
git_called["cmd"] = cmd
return SimpleNamespace(returncode=0, stderr="", stdout="")
monkeypatch.setattr(module.subprocess, "run", _fake_run)
cleaned = module.cleanup_artifacts(8, confirm=True)
assert not draft_path.exists()
assert git_called["count"] == 1
assert git_called["cmd"] == ["git", "reset", "HEAD", "."]
assert any("Git 暂存区已清理" in item for item in cleaned)
backup_dir = tmp_path / ".noma" / "recovery_backups"
backups = list(backup_dir.glob("ch0008-*"))
assert backups
+968
View File
@@ -0,0 +1,968 @@
#!/usr/bin/env python3
# -*- coding: utf-8 -*-
"""
wiki_manager — 结构化 Wiki/Notebook 知识管理器
职责:
- 实体档案维护(从 index.db 全量同步到 .noma/wiki/entities/)
- 伏笔/剧情线索管理(从 state.json 同步到 .noma/wiki/plot/)
- 关系图谱维护(从 index.db relationships 同步到 .noma/wiki/relationships/)
- 写作模式记录(替代 project_memory.json 的死胡同,写入 .noma/wiki/patterns/)
- 纯 grep 搜索(无 embedding 依赖)
设计原则:
- Wiki = 地面真相(ground truth),从 index.db/state.json 全量重写
- RAG = 语义检索(fuzzy context),负责向量/BM25 搜索
- 两者互补:Wiki 提供确定性事实,RAG 提供语义相关上下文
"""
from __future__ import annotations
import argparse
import json
import logging
import re
import sys
from datetime import datetime, timezone
from pathlib import Path
from typing import Any, Dict, List, Optional
from runtime_compat import normalize_windows_path
logger = logging.getLogger(__name__)
# ---------------------------------------------------------------------------
# YAML frontmatter parser (lightweight, no PyYAML dependency)
# ---------------------------------------------------------------------------
_FRONTMATTER_RE = re.compile(r"^---\s*\n(.*?)\n---\s*\n", re.DOTALL)
def _parse_frontmatter(text: str) -> tuple[Dict[str, Any], str]:
"""Parse YAML frontmatter from markdown text.
Returns (frontmatter_dict, body_without_frontmatter).
Only supports simple key: value and key: [list] syntax.
"""
m = _FRONTMATTER_RE.match(text)
if not m:
return {}, text
fm_text = m.group(1)
body = text[m.end():]
result: Dict[str, Any] = {}
for line in fm_text.split("\n"):
line = line.strip()
if not line or line.startswith("#"):
continue
if ":" not in line:
continue
key, _, val = line.partition(":")
key = key.strip()
val = val.strip()
if not key:
continue
# Handle list values: [a, b, c]
if val.startswith("[") and val.endswith("]"):
items = [x.strip().strip("\"'") for x in val[1:-1].split(",") if x.strip()]
result[key] = items
# Handle quoted strings
elif (val.startswith('"') and val.endswith('"')) or (
val.startswith("'") and val.endswith("'")
):
result[key] = val[1:-1]
# Handle booleans
elif val.lower() in ("true", "yes"):
result[key] = True
elif val.lower() in ("false", "no"):
result[key] = False
# Handle numbers
elif val.isdigit():
result[key] = int(val)
else:
try:
result[key] = float(val)
except ValueError:
result[key] = val
return result, body
def _serialize_frontmatter(data: Dict[str, Any]) -> str:
"""Serialize dict to YAML frontmatter string."""
lines = ["---"]
for key, val in data.items():
if isinstance(val, list):
items = ", ".join(str(v) for v in val)
lines.append(f"{key}: [{items}]")
elif isinstance(val, bool):
lines.append(f"{key}: {'true' if val else 'false'}")
elif isinstance(val, (int, float)):
lines.append(f"{key}: {val}")
elif val is None:
lines.append(f"{key}:")
else:
lines.append(f"{key}: \"{val}\"")
lines.append("---")
return "\n".join(lines)
# ---------------------------------------------------------------------------
# WikiManager
# ---------------------------------------------------------------------------
class WikiManager:
"""Wiki/Notebook manager for structured knowledge storage."""
def __init__(self, config: Any = None):
if config is None:
from .config import get_config
config = get_config()
self.config = config
@property
def wiki_dir(self) -> Path:
return self.config.noma_dir / "wiki"
def ensure_wiki_dirs(self) -> None:
"""Create wiki directory structure if it doesn't exist."""
for subdir in ("entities", "plot", "relationships", "patterns"):
(self.wiki_dir / subdir).mkdir(parents=True, exist_ok=True)
def _now_iso(self) -> str:
return datetime.now(timezone.utc).isoformat(timespec="seconds")
# -----------------------------------------------------------------------
# Entity Wiki
# -----------------------------------------------------------------------
def update_entity_wiki(
self,
entity_id: str,
entity_data: Dict[str, Any],
state_changes: Optional[List[Dict[str, Any]]] = None,
aliases: Optional[List[str]] = None,
relationships: Optional[List[Dict[str, Any]]] = None,
) -> Path:
"""Create or update an entity wiki file from index.db data.
Performs a full rewrite (wiki = latest state snapshot).
"""
self.ensure_wiki_dirs()
entity_id = str(entity_id or "").strip()
if not entity_id:
raise ValueError("entity_id is required")
canonical_name = str(entity_data.get("canonical_name") or entity_id)
entity_type = str(entity_data.get("type") or "未知")
tier = str(entity_data.get("tier") or "装饰")
desc = str(entity_data.get("desc") or "")
current = entity_data.get("current") or {}
if isinstance(current, str):
try:
current = json.loads(current)
except (json.JSONDecodeError, TypeError):
current = {}
first_appearance = entity_data.get("first_appearance") or 0
last_appearance = entity_data.get("last_appearance") or 0
is_protagonist = bool(entity_data.get("is_protagonist"))
# Build frontmatter
frontmatter: Dict[str, Any] = {
"id": entity_id,
"type": entity_type,
"canonical_name": canonical_name,
"tier": tier,
"first_appearance": first_appearance,
"last_appearance": last_appearance,
"updated_at": self._now_iso(),
}
if is_protagonist:
frontmatter["is_protagonist"] = True
# Build body
lines: List[str] = []
lines.append(f"# {canonical_name}")
lines.append("")
# Basic info
lines.append("## 基本信息")
lines.append(f"- **类型**: {entity_type} / {tier}")
if aliases:
lines.append(f"- **别名**: {', '.join(aliases)}")
lines.append(f"- **首次出场**: 第{first_appearance}章")
lines.append(f"- **最近出场**: 第{last_appearance}章")
if desc:
lines.append(f"- **描述**: {desc}")
lines.append("")
# Current state
if current:
lines.append("## 当前状态")
for k, v in current.items():
if isinstance(v, dict):
lines.append(f"- **{k}**:")
for sk, sv in v.items():
lines.append(f" - {sk}: {sv}")
elif isinstance(v, list):
lines.append(f"- **{k}**: {', '.join(str(x) for x in v)}")
else:
lines.append(f"- **{k}**: {v}")
lines.append("")
# Relationships
if relationships:
lines.append("## 关系")
for rel in relationships:
from_e = str(rel.get("from_entity") or rel.get("from") or "")
to_e = str(rel.get("to_entity") or rel.get("to") or "")
rel_type = str(rel.get("type") or "关联")
desc_text = str(rel.get("description") or "")
other = to_e if from_e == entity_id else from_e
suffix = f" ({desc_text})" if desc_text else ""
lines.append(f"- {other}: {rel_type}{suffix}")
lines.append("")
# State change history
if state_changes:
lines.append("## 状态变化历史")
lines.append("| 章节 | 字段 | 旧值 | 新值 | 原因 |")
lines.append("|------|------|------|------|------|")
for sc in state_changes[:50]: # Cap at 50 rows
ch = sc.get("chapter", "?")
field = sc.get("field", "?")
old = sc.get("old_value", "")
new = sc.get("new_value", "")
reason = sc.get("reason", "")
lines.append(f"| {ch} | {field} | {old} | {new} | {reason} |")
lines.append("")
# Write file
content = _serialize_frontmatter(frontmatter) + "\n\n" + "\n".join(lines)
file_path = self.wiki_dir / "entities" / f"{entity_id}.md"
file_path.parent.mkdir(parents=True, exist_ok=True)
file_path.write_text(content, encoding="utf-8")
return file_path
def get_entity_wiki(self, entity_id: str) -> Optional[Dict[str, Any]]:
"""Read an entity wiki file and return parsed frontmatter + body."""
file_path = self.wiki_dir / "entities" / f"{entity_id}.md"
if not file_path.exists():
return None
text = file_path.read_text(encoding="utf-8")
fm, body = _parse_frontmatter(text)
return {"frontmatter": fm, "body": body, "path": str(file_path)}
def list_entity_wiki(
self,
entity_type: Optional[str] = None,
tier: Optional[str] = None,
) -> List[Dict[str, Any]]:
"""List entity wiki entries with optional filters."""
entities_dir = self.wiki_dir / "entities"
if not entities_dir.exists():
return []
results: List[Dict[str, Any]] = []
for f in sorted(entities_dir.glob("*.md")):
text = f.read_text(encoding="utf-8")
fm, _ = _parse_frontmatter(text)
if entity_type and fm.get("type") != entity_type:
continue
if tier and fm.get("tier") != tier:
continue
results.append({
"id": fm.get("id", f.stem),
"canonical_name": fm.get("canonical_name", f.stem),
"type": fm.get("type", "未知"),
"tier": fm.get("tier", "装饰"),
"first_appearance": fm.get("first_appearance", 0),
"last_appearance": fm.get("last_appearance", 0),
"path": str(f),
})
return results
# -----------------------------------------------------------------------
# Plot Wiki
# -----------------------------------------------------------------------
def update_plot_threads(
self,
foreshadowing: Optional[List[Dict[str, Any]]] = None,
constraints: Optional[Dict[str, Any]] = None,
strand_tracker: Optional[Dict[str, Any]] = None,
) -> Path:
"""Update the plot/threads.md file."""
self.ensure_wiki_dirs()
frontmatter: Dict[str, Any] = {
"type": "plot_threads",
"updated_at": self._now_iso(),
}
lines: List[str] = []
lines.append("# 伏笔与剧情线索")
lines.append("")
# Foreshadowing (may be list of dicts or list of strings)
if foreshadowing:
normalized: List[Dict[str, Any]] = []
for f in foreshadowing:
if isinstance(f, str):
normalized.append({"content": f, "status": "进行中"})
elif isinstance(f, dict):
normalized.append(f)
active = [f for f in normalized if f.get("status") != "已回收"]
resolved = [f for f in normalized if f.get("status") == "已回收"]
if active:
lines.append("## 活跃伏笔")
lines.append("")
for i, ft in enumerate(active, 1):
fid = ft.get("id", f"FT-{i:03d}")
title = ft.get("title") or ft.get("content", "未命名")
planted = ft.get("planted_chapter") or ft.get("chapter", "?")
target = ft.get("target_chapter", "?")
tier_val = ft.get("tier", "支线")
content = ft.get("content", "")
lines.append(f"### {fid}: {title}")
lines.append(f"- **埋设章节**: 第{planted}章")
lines.append(f"- **目标章节**: 第{target}章")
lines.append(f"- **状态**: 进行中")
lines.append(f"- **层级**: {tier_val}")
if content:
lines.append(f"- **内容**: {content}")
lines.append("")
if resolved:
lines.append("## 已回收伏笔")
lines.append("")
for ft in resolved:
title = ft.get("title") or ft.get("content", "未命名")
planted = ft.get("planted_chapter") or ft.get("chapter", "?")
lines.append(f"- 第{planted}章: {title}")
lines.append("")
# Strand tracker
if strand_tracker:
lines.append("## 节奏追踪 (Strand Weave)")
lines.append(f"- **当前主导**: {strand_tracker.get('current_dominant', 'quest')}")
lines.append(f"- **距上次切换**: {strand_tracker.get('chapters_since_switch', 0)}章")
last_q = strand_tracker.get("last_quest_chapter", 0)
last_f = strand_tracker.get("last_fire_chapter", 0)
last_c = strand_tracker.get("last_constellation_chapter", 0)
lines.append(f"- **最近Quest**: 第{last_q}章")
lines.append(f"- **最近Fire**: 第{last_f}章")
lines.append(f"- **最近Constellation**: 第{last_c}章")
lines.append("")
# Constraints
if constraints:
lines.append("## 创作约束")
if constraints.get("anti_trope"):
lines.append(f"- **反套路**: {constraints['anti_trope']}")
if constraints.get("hard_constraints"):
for hc in constraints["hard_constraints"]:
lines.append(f"- **硬约束**: {hc}")
if constraints.get("protagonist_flaw"):
lines.append(f"- **主角缺陷**: {constraints['protagonist_flaw']}")
if constraints.get("antagonist_mirror"):
lines.append(f"- **反派镜像**: {constraints['antagonist_mirror']}")
lines.append("")
content = _serialize_frontmatter(frontmatter) + "\n\n" + "\n".join(lines)
file_path = self.wiki_dir / "plot" / "threads.md"
file_path.parent.mkdir(parents=True, exist_ok=True)
file_path.write_text(content, encoding="utf-8")
return file_path
def get_plot_threads(self) -> Optional[Dict[str, Any]]:
"""Read and parse plot threads wiki."""
file_path = self.wiki_dir / "plot" / "threads.md"
if not file_path.exists():
return None
text = file_path.read_text(encoding="utf-8")
fm, body = _parse_frontmatter(text)
return {"frontmatter": fm, "body": body, "path": str(file_path)}
# -----------------------------------------------------------------------
# Relationship Wiki
# -----------------------------------------------------------------------
def update_relationship_graph(
self,
relationships: List[Dict[str, Any]],
entity_names: Optional[Dict[str, str]] = None,
) -> Path:
"""Update relationships/graph.md from index.db relationships table."""
self.ensure_wiki_dirs()
entity_names = entity_names or {}
frontmatter: Dict[str, Any] = {
"type": "relationship_graph",
"updated_at": self._now_iso(),
"edge_count": len(relationships),
}
lines: List[str] = []
lines.append("# 关系图谱")
lines.append("")
if not relationships:
lines.append("暂无关系数据。")
else:
# Group by entity
by_entity: Dict[str, List[Dict[str, Any]]] = {}
for rel in relationships:
from_e = str(rel.get("from_entity") or rel.get("from") or "")
to_e = str(rel.get("to_entity") or rel.get("to") or "")
if from_e:
by_entity.setdefault(from_e, []).append(rel)
if to_e:
by_entity.setdefault(to_e, []).append(rel)
for entity_id in sorted(by_entity.keys()):
name = entity_names.get(entity_id, entity_id)
lines.append(f"## {name}")
lines.append("")
for rel in by_entity[entity_id]:
from_e = str(rel.get("from_entity") or rel.get("from") or "")
to_e = str(rel.get("to_entity") or rel.get("to") or "")
rel_type = str(rel.get("type") or "关联")
desc = str(rel.get("description") or "")
ch = rel.get("chapter", "?")
other_name = entity_names.get(to_e if from_e == entity_id else from_e, to_e if from_e == entity_id else from_e)
suffix = f" — {desc}" if desc else ""
lines.append(f"- {other_name}: {rel_type} (第{ch}章){suffix}")
lines.append("")
content = _serialize_frontmatter(frontmatter) + "\n\n" + "\n".join(lines)
file_path = self.wiki_dir / "relationships" / "graph.md"
file_path.parent.mkdir(parents=True, exist_ok=True)
file_path.write_text(content, encoding="utf-8")
return file_path
# -----------------------------------------------------------------------
# Writing Patterns (replaces project_memory.json dead-end)
# -----------------------------------------------------------------------
def append_writing_pattern(
self,
pattern_type: str,
description: str,
source_chapter: int,
details: Optional[str] = None,
) -> Path:
"""Append a writing pattern to wiki/patterns/writing-patterns.md."""
self.ensure_wiki_dirs()
file_path = self.wiki_dir / "patterns" / "writing-patterns.md"
# Read existing or initialize
if file_path.exists():
text = file_path.read_text(encoding="utf-8")
fm, body = _parse_frontmatter(text)
else:
fm = {"type": "writing_patterns"}
body = "# 写作模式库\n\n## 模式列表\n"
# Update frontmatter
fm["updated_at"] = self._now_iso()
pattern_count = fm.get("pattern_count", 0) + 1
fm["pattern_count"] = pattern_count
# Append new pattern
now = self._now_iso()
body = body.rstrip() + "\n\n"
body += f"### P-{pattern_count:03d}\n"
body += f"- **类型**: {pattern_type}\n"
body += f"- **描述**: {description}\n"
body += f"- **来源章节**: 第{source_chapter}章\n"
body += f"- **记录时间**: {now}\n"
if details:
body += f"- **详情**: {details}\n"
content = _serialize_frontmatter(fm) + "\n\n" + body
file_path.parent.mkdir(parents=True, exist_ok=True)
file_path.write_text(content, encoding="utf-8")
return file_path
def get_writing_patterns(
self, pattern_type: Optional[str] = None
) -> List[Dict[str, Any]]:
"""Read writing patterns, optionally filtered by type."""
file_path = self.wiki_dir / "patterns" / "writing-patterns.md"
if not file_path.exists():
return []
text = file_path.read_text(encoding="utf-8")
_, body = _parse_frontmatter(text)
patterns: List[Dict[str, Any]] = []
current_pattern: Dict[str, Any] = {}
for line in body.split("\n"):
line = line.strip()
if line.startswith("### P-"):
if current_pattern:
patterns.append(current_pattern)
current_pattern = {"id": line[4:]}
elif line.startswith("- **类型**:"):
current_pattern["pattern_type"] = line.split(":", 1)[1].strip()
elif line.startswith("- **描述**:"):
current_pattern["description"] = line.split(":", 1)[1].strip()
elif line.startswith("- **来源章节**:"):
ch_text = line.split(":", 1)[1].strip()
ch_match = re.search(r"\d+", ch_text)
current_pattern["source_chapter"] = int(ch_match.group()) if ch_match else 0
elif line.startswith("- **记录时间**:"):
current_pattern["learned_at"] = line.split(":", 1)[1].strip()
elif line.startswith("- **详情**:"):
current_pattern["details"] = line.split(":", 1)[1].strip()
if current_pattern:
patterns.append(current_pattern)
if pattern_type:
patterns = [p for p in patterns if p.get("pattern_type") == pattern_type]
return patterns
# -----------------------------------------------------------------------
# Search (pure grep)
# -----------------------------------------------------------------------
def search_wiki(
self, query: str, wiki_type: Optional[str] = None
) -> List[Dict[str, Any]]:
"""Grep-based search across all wiki files.
Args:
query: Search text (supports Chinese and English)
wiki_type: Optional filter: "entity", "plot", "relationship", "pattern"
"""
if not query or not query.strip():
return []
self.ensure_wiki_dirs()
query_lower = query.lower().strip()
type_dirs = {
"entity": "entities",
"plot": "plot",
"relationship": "relationships",
"pattern": "patterns",
}
search_dirs: List[Path] = []
if wiki_type and wiki_type in type_dirs:
search_dirs.append(self.wiki_dir / type_dirs[wiki_type])
else:
for subdir in type_dirs.values():
search_dirs.append(self.wiki_dir / subdir)
results: List[Dict[str, Any]] = []
for search_dir in search_dirs:
if not search_dir.exists():
continue
for md_file in sorted(search_dir.glob("*.md")):
try:
text = md_file.read_text(encoding="utf-8")
except Exception:
continue
if query_lower not in text.lower():
continue
fm, body = _parse_frontmatter(text)
# Find matching lines
matching_lines: List[str] = []
for line in text.split("\n"):
if query_lower in line.lower():
matching_lines.append(line.strip())
results.append({
"file": str(md_file.relative_to(self.wiki_dir)),
"type": fm.get("type", md_file.parent.name),
"id": fm.get("id", md_file.stem),
"name": fm.get("canonical_name", fm.get("id", md_file.stem)),
"matches": matching_lines[:5], # Cap at 5 matching lines
"match_count": len(matching_lines),
})
return results
# -----------------------------------------------------------------------
# Sync from index.db / state.json
# -----------------------------------------------------------------------
def sync_from_index(
self, entity_ids: Optional[List[str]] = None
) -> Dict[str, Any]:
"""Bulk sync wiki entries from index.db.
If entity_ids is None, sync all non-archived entities.
"""
from .index_manager import IndexManager
idx = IndexManager(self.config)
self.ensure_wiki_dirs()
if entity_ids:
entities = []
for eid in entity_ids:
e = idx.get_entity(eid)
if e:
entities.append(e)
else:
entities = idx.get_core_entities()
synced = 0
errors: List[str] = []
for entity in entities:
try:
eid = entity.get("id", "")
if not eid:
continue
aliases = idx.get_entity_aliases(eid)
relationships = idx.get_entity_relationships(eid, direction="both")
state_changes = idx.get_entity_state_changes(eid, limit=50)
self.update_entity_wiki(
entity_id=eid,
entity_data=entity,
state_changes=state_changes,
aliases=aliases,
relationships=relationships,
)
synced += 1
except Exception as exc:
errors.append(f"{entity.get('id', '?')}: {exc}")
logger.warning("wiki sync error for entity %s: %s", entity.get("id"), exc)
# Sync relationship graph
try:
all_relationships = idx.get_recent_relationships(limit=500)
entity_names = {
e["id"]: e.get("canonical_name", e["id"])
for e in entities
}
self.update_relationship_graph(all_relationships, entity_names=entity_names)
except Exception as exc:
errors.append(f"relationships: {exc}")
logger.warning("wiki sync error for relationships: %s", exc)
return {"synced": synced, "total": len(entities), "errors": errors}
def sync_from_state(self) -> Dict[str, Any]:
"""Sync plot threads and constraints from state.json + genesis_contract.json."""
self.ensure_wiki_dirs()
state_file = self.config.state_file
if not state_file.exists():
return {"error": "state.json not found"}
state = json.loads(state_file.read_text(encoding="utf-8"))
# Foreshadowing
plot_threads = state.get("plot_threads", {})
foreshadowing = plot_threads.get("foreshadowing", [])
# Strand tracker
strand_tracker = state.get("strand_tracker", {})
# Genesis contract (constraints)
genesis_path = self.config.noma_dir / "genesis_contract.json"
constraints: Dict[str, Any] = {}
if genesis_path.exists():
try:
genesis = json.loads(genesis_path.read_text(encoding="utf-8"))
core_desire = genesis.get("core_desire", {})
constraints = {
"anti_trope": genesis.get("anti_trope", ""),
"hard_constraints": core_desire.get("taboos", []),
"protagonist_flaw": core_desire.get("flaw", ""),
"antagonist_mirror": genesis.get("antagonist_mirror", ""),
}
except Exception as exc:
logger.warning("failed to read genesis_contract.json: %s", exc)
# Idea bank
idea_bank_path = self.config.noma_dir / "novel_data" / "idea_bank.json"
if idea_bank_path.exists():
try:
idea_bank = json.loads(idea_bank_path.read_text(encoding="utf-8"))
inherited = idea_bank.get("constraints_inherited", {})
if not constraints.get("anti_trope"):
constraints["anti_trope"] = inherited.get("anti_trope", "")
if not constraints.get("hard_constraints"):
constraints["hard_constraints"] = inherited.get("hard_constraints", [])
if not constraints.get("protagonist_flaw"):
constraints["protagonist_flaw"] = inherited.get("protagonist_flaw", "")
except Exception as exc:
logger.warning("failed to read idea_bank.json: %s", exc)
self.update_plot_threads(
foreshadowing=foreshadowing,
constraints=constraints or None,
strand_tracker=strand_tracker or None,
)
return {
"foreshadowing_count": len(foreshadowing),
"has_constraints": bool(constraints),
"has_strand_tracker": bool(strand_tracker),
}
# -----------------------------------------------------------------------
# Migration
# -----------------------------------------------------------------------
def migrate_from_project_memory(self) -> Dict[str, Any]:
"""Migrate existing project_memory.json patterns to wiki."""
self.ensure_wiki_dirs()
# Look for project_memory.json in various locations
candidates = [
self.config.noma_dir / "novel_data" / "project_memory.json",
self.config.project_root / "novelmaster" / "project_memory.json",
self.config.project_root / "project_memory.json",
]
migrated = 0
for pm_path in candidates:
if not pm_path.exists():
continue
try:
data = json.loads(pm_path.read_text(encoding="utf-8"))
patterns = data.get("patterns", [])
for p in patterns:
self.append_writing_pattern(
pattern_type=p.get("pattern_type", "unknown"),
description=p.get("description", ""),
source_chapter=p.get("source_chapter", 0),
)
migrated += 1
if migrated > 0:
logger.info("migrated %d patterns from %s", migrated, pm_path)
break
except Exception as exc:
logger.warning("failed to migrate from %s: %s", pm_path, exc)
return {"migrated": migrated}
def rebuild_index(self) -> Path:
"""Regenerate _index.md from all wiki files."""
self.ensure_wiki_dirs()
lines: List[str] = []
lines.append("# Wiki 索引")
lines.append("")
lines.append(f"更新时间: {self._now_iso()}")
lines.append("")
# Entities
entities = self.list_entity_wiki()
lines.append(f"## 实体 ({len(entities)})")
lines.append("")
for e in entities:
lines.append(f"- [{e['canonical_name']}](entities/{e['id']}.md) — {e['type']} / {e['tier']}")
lines.append("")
# Plot
plot = self.get_plot_threads()
if plot:
lines.append("## 伏笔与剧情线索")
lines.append(f"- [threads.md](plot/threads.md)")
lines.append("")
# Relationships
rel_path = self.wiki_dir / "relationships" / "graph.md"
if rel_path.exists():
lines.append("## 关系图谱")
lines.append(f"- [graph.md](relationships/graph.md)")
lines.append("")
# Patterns
patterns = self.get_writing_patterns()
lines.append(f"## 写作模式 ({len(patterns)})")
lines.append(f"- [writing-patterns.md](patterns/writing-patterns.md)")
lines.append("")
index_path = self.wiki_dir / "_index.md"
index_path.write_text("\n".join(lines), encoding="utf-8")
return index_path
# ---------------------------------------------------------------------------
# CLI interface
# ---------------------------------------------------------------------------
def _parse_args(argv: list[str]) -> argparse.Namespace:
parser = argparse.ArgumentParser(description="wiki_manager CLI")
parser.add_argument("--project-root", required=True, help="项目根目录")
sub = parser.add_subparsers(dest="command", required=True)
# update-entity
p_ue = sub.add_parser("update-entity", help="同步单个实体到 wiki")
p_ue.add_argument("--id", required=True, help="实体 ID")
# update-plot
sub.add_parser("update-plot", help="同步伏笔/剧情线索到 wiki")
# update-relationship
sub.add_parser("update-relationship", help="同步关系图谱到 wiki")
# update-patterns
p_up = sub.add_parser("update-patterns", help="添加写作模式")
p_up.add_argument("--data", required=True, help="JSON 格式模式数据")
# search
p_s = sub.add_parser("search", help="搜索 wiki")
p_s.add_argument("--query", required=True, help="搜索关键词")
p_s.add_argument("--type", dest="wiki_type", help="类型过滤: entity|plot|relationship|pattern")
# sync-from-index
p_sfi = sub.add_parser("sync-from-index", help="从 index.db 批量同步实体")
p_sfi.add_argument("--entity-ids", help="JSON 格式实体 ID 列表(可选,默认同步所有核心实体)")
# sync-from-state
sub.add_parser("sync-from-state", help="从 state.json 同步伏笔/约束")
# migrate-project-memory
sub.add_parser("migrate-project-memory", help="迁移 project_memory.json 到 wiki")
# rebuild-index
sub.add_parser("rebuild-index", help="重建 _index.md")
# list
p_l = sub.add_parser("list", help="列出 wiki 条目")
p_l.add_argument("--type", dest="wiki_type", help="类型过滤: entity|plot|relationship|pattern")
# get
p_g = sub.add_parser("get", help="获取单个 wiki 条目")
p_g.add_argument("--id", required=True, help="条目 ID")
p_g.add_argument("--type", dest="wiki_type", default="entity", help="类型: entity|plot|pattern")
return parser.parse_args(argv)
def main() -> None:
args = _parse_args(sys.argv[1:])
from .config import DataModulesConfig
config = DataModulesConfig.from_project_root(args.project_root)
wiki = WikiManager(config)
if args.command == "update-entity":
from .index_manager import IndexManager
idx = IndexManager(config)
entity = idx.get_entity(args.id)
if not entity:
print(json.dumps({"error": f"entity not found: {args.id}"}, ensure_ascii=False))
raise SystemExit(1)
aliases = idx.get_entity_aliases(args.id)
relationships = idx.get_entity_relationships(args.id, direction="both")
state_changes = idx.get_entity_state_changes(args.id, limit=50)
path = wiki.update_entity_wiki(
entity_id=args.id,
entity_data=entity,
state_changes=state_changes,
aliases=aliases,
relationships=relationships,
)
print(json.dumps({"ok": True, "path": str(path)}, ensure_ascii=False))
elif args.command == "update-plot":
result = wiki.sync_from_state()
print(json.dumps({"ok": True, **result}, ensure_ascii=False))
elif args.command == "update-relationship":
from .index_manager import IndexManager
idx = IndexManager(config)
relationships = idx.get_recent_relationships(limit=500)
entities = idx.get_core_entities()
entity_names = {e["id"]: e.get("canonical_name", e["id"]) for e in entities}
path = wiki.update_relationship_graph(relationships, entity_names=entity_names)
print(json.dumps({"ok": True, "path": str(path)}, ensure_ascii=False))
elif args.command == "update-patterns":
data = json.loads(args.data)
path = wiki.append_writing_pattern(
pattern_type=data.get("pattern_type", "unknown"),
description=data.get("description", ""),
source_chapter=data.get("source_chapter", 0),
details=data.get("details"),
)
print(json.dumps({"ok": True, "path": str(path)}, ensure_ascii=False))
elif args.command == "search":
results = wiki.search_wiki(args.query, wiki_type=args.wiki_type)
print(json.dumps({"results": results, "count": len(results)}, ensure_ascii=False, indent=2))
elif args.command == "sync-from-index":
entity_ids = None
if args.entity_ids:
entity_ids = json.loads(args.entity_ids)
result = wiki.sync_from_index(entity_ids=entity_ids)
print(json.dumps({"ok": True, **result}, ensure_ascii=False))
elif args.command == "sync-from-state":
result = wiki.sync_from_state()
print(json.dumps({"ok": True, **result}, ensure_ascii=False))
elif args.command == "migrate-project-memory":
result = wiki.migrate_from_project_memory()
print(json.dumps({"ok": True, **result}, ensure_ascii=False))
elif args.command == "rebuild-index":
path = wiki.rebuild_index()
print(json.dumps({"ok": True, "path": str(path)}, ensure_ascii=False))
elif args.command == "list":
if args.wiki_type == "entity" or not args.wiki_type:
entities = wiki.list_entity_wiki()
for e in entities:
print(json.dumps(e, ensure_ascii=False))
if args.wiki_type == "pattern" or not args.wiki_type:
patterns = wiki.get_writing_patterns()
for p in patterns:
print(json.dumps(p, ensure_ascii=False))
elif args.command == "get":
if args.wiki_type == "entity":
result = wiki.get_entity_wiki(args.id)
elif args.wiki_type == "plot":
result = wiki.get_plot_threads()
elif args.wiki_type == "pattern":
patterns = wiki.get_writing_patterns()
result = next((p for p in patterns if p.get("id") == args.id), None)
else:
result = None
if result:
print(json.dumps(result, ensure_ascii=False, indent=2))
else:
print(json.dumps({"error": "not found"}, ensure_ascii=False))
raise SystemExit(1)
if __name__ == "__main__":
main()
@@ -0,0 +1,478 @@
#!/usr/bin/env python3
# -*- coding: utf-8 -*-
"""
Writing guidance and checklist builders.
"""
from __future__ import annotations
from typing import Any, Dict, List
from .genre_aliases import to_profile_key
GENRE_GUIDANCE_TEXT: dict[str, str] = {
"xianxia": "题材加权:强化升级/对抗结果的可见反馈,术语解释后置。",
"shuangwen": "题材加权:维持高爽点密度,主爽点外叠加一个副轴反差。",
"urban-power": "题材加权:优先写社会反馈链(他人反应→资源变化→地位变化)。",
"romance": "题材加权:每章推进关系位移,避免情绪原地打转。",
"mystery": "题材加权:线索必须可回收,优先以规则冲突制造悬念。",
"rules-mystery": "题材加权:规则先于解释,代价先于胜利。",
"zhihu-short": "题材加权:压缩铺垫,优先反转与高强度结尾钩。",
"substitute": "题材加权:强化误解-拉扯-决断链路,避免重复虐点。",
"esports": "题材加权:每场对抗至少写清一个战术决策点与其后果。",
"livestream": "题材加权:强化“外部反馈→主角反制→数据变化”即时闭环。",
"cosmic-horror": "题材加权:恐怖来源于规则与代价,不依赖空泛惊悚形容。",
}
GENRE_METHOD_ANCHORS: dict[str, dict[str, str]] = {
"xianxia": {
"pressure_source": "资源争夺/境界压制",
"release_target": "主角主动破局并拿到可见收益",
},
"urban-power": {
"pressure_source": "阶层卡位/权力压制",
"release_target": "主角通过资源博弈拿到地位与回报",
},
"romance": {
"pressure_source": "关系误解/情感拉扯",
"release_target": "关系位移落地并形成下一步承诺",
},
"mystery": {
"pressure_source": "线索缺失/规则冲突",
"release_target": "给出可验证的新线索并保留未知区",
},
"rules-mystery": {
"pressure_source": "规则反噬/代价递增",
"release_target": "用代价换突破并留下更高阶规则问题",
},
"zhihu-short": {
"pressure_source": "信息落差/立场对撞",
"release_target": "反转兑现并形成高强度尾钩",
},
"substitute": {
"pressure_source": "身份误读/情绪对峙",
"release_target": "误解链推进到明确决断",
},
"esports": {
"pressure_source": "战术压制/节奏失衡",
"release_target": "关键决策生效并转化为局势优势",
},
"livestream": {
"pressure_source": "舆论波动/数据下滑",
"release_target": "当场反制形成可见数据回弹",
},
"cosmic-horror": {
"pressure_source": "认知失真/规则侵蚀",
"release_target": "以明确代价换阶段性生存窗口",
},
"history-travel": {
"pressure_source": "历史惯性/礼教阻力",
"release_target": "知识优势兑现并引发新的连锁反应",
},
"game-lit": {
"pressure_source": "系统规则限制/资源稀缺",
"release_target": "数值突破并暴露更高层级威胁",
},
}
def build_methodology_strategy_card(
*,
chapter: int,
reader_signal: Dict[str, Any],
genre_profile: Dict[str, Any],
label: str = "digital-serial-v1",
) -> Dict[str, Any]:
genre = str(genre_profile.get("genre") or "").strip()
profile_key = to_profile_key(genre) or "general"
hook_usage = reader_signal.get("hook_type_usage") or {}
pattern_usage = reader_signal.get("pattern_usage") or {}
review_trend = reader_signal.get("review_trend") or {}
low_ranges = reader_signal.get("low_score_ranges") or []
dominant_hook = ""
if isinstance(hook_usage, dict) and hook_usage:
dominant_hook = max(hook_usage.items(), key=lambda kv: kv[1])[0]
dominant_pattern = ""
if isinstance(pattern_usage, dict) and pattern_usage:
dominant_pattern = max(pattern_usage.items(), key=lambda kv: kv[1])[0]
overall_avg = float(review_trend.get("overall_avg") or 0.0)
has_low_range = bool(low_ranges)
hook_variety = len(hook_usage) if isinstance(hook_usage, dict) else 0
pattern_variety = len(pattern_usage) if isinstance(pattern_usage, dict) else 0
next_reason_clarity = 70.0 + (4.0 if has_low_range else 8.0)
anchor_effectiveness = 68.0 + (6.0 if dominant_hook else 0.0) + (4.0 if overall_avg >= 75 else -4.0)
rhythm_naturalness = 65.0 + min(10.0, float(hook_variety + pattern_variety) * 2.0)
risk_flags: List[str] = []
if has_low_range:
risk_flags.append("low_score_recency")
if dominant_pattern:
risk_flags.append("pattern_overuse_watch")
if overall_avg > 0 and overall_avg < 75:
risk_flags.append("readability_guard")
stage_mod = chapter % 5
if stage_mod in {1, 2}:
stage = "build_up"
elif stage_mod in {3, 4}:
stage = "confront"
else:
stage = "release"
anchor_preset = GENRE_METHOD_ANCHORS.get(
profile_key,
{
"pressure_source": "生存目标/资源竞争",
"release_target": "主角完成阶段目标并留下新的行动理由",
},
)
return {
"enabled": True,
"framework": label,
"pilot": profile_key,
"genre_profile_key": profile_key,
"chapter_stage": stage,
"emotion_anchor": {
"pressure_source": anchor_preset["pressure_source"],
"release_target": anchor_preset["release_target"],
"position_hint": "前段设压,中后段释放,避免固定字位打点",
},
"long_arc_controls": {
"map_transition": "阶段切换承接既有资产与关系账本,避免能力与收益归零",
"power_guard": "关键胜利必须给机制理由(信息/资源/代价/策略)",
"antagonist_model": "反派需具备目标-手段-代价三要素,避免工具人推进",
},
"serialization_ops": {
"next_reason": "章末或后段给出可复述的下一章动机句",
"interaction_note": "保留一个可讨论分歧点,便于连载互动反馈",
},
"observability": {
"next_reason_clarity": round(max(0.0, min(100.0, next_reason_clarity)), 2),
"anchor_effectiveness": round(max(0.0, min(100.0, anchor_effectiveness)), 2),
"rhythm_naturalness": round(max(0.0, min(100.0, rhythm_naturalness)), 2),
},
"signals": {
"dominant_hook": dominant_hook,
"dominant_pattern": dominant_pattern,
"risk_flags": risk_flags,
},
}
def build_methodology_guidance_items(strategy_card: Dict[str, Any]) -> List[str]:
if not isinstance(strategy_card, dict) or not strategy_card.get("enabled"):
return []
observability = strategy_card.get("observability") or {}
signals = strategy_card.get("signals") or {}
risk_flags = list(signals.get("risk_flags") or [])
stage = str(strategy_card.get("chapter_stage") or "build_up")
genre_key = str(strategy_card.get("genre_profile_key") or strategy_card.get("pilot") or "general")
stage_text = {
"build_up": "本章以铺压为主,优先做威胁与代价的可感知铺垫。",
"confront": "本章以正面对抗为主,确保破局路径清晰可复盘。",
"release": "本章以释放与余波为主,给出实质收益并引出下一问。",
}.get(stage, "本章保持压力-破局-余波的完整链路。")
items = [
f"方法论策略(通用/{genre_key}):{stage_text}",
"长线控制:换图承接旧资产,避免主角进入新地图后能力与资源归零。",
"机制控制:关键胜利必须写出机制理由与代价,不用纯光环碾压。",
(
"连载互动:保留一个可讨论分歧点,强化下章追更动机。"
f"(next_reason={observability.get('next_reason_clarity')})"
),
]
if "pattern_overuse_watch" in risk_flags:
dominant_pattern = str(signals.get("dominant_pattern") or "").strip()
if dominant_pattern:
items.append(f"风险修正:近期“{dominant_pattern}”偏高频,本章补一个异质副轴避免疲劳。")
if "readability_guard" in risk_flags:
items.append("风险修正:近期审查均分偏低,本章优先保证段落动作-结果闭环与可读性。")
return items
def build_guidance_items(
*,
chapter: int,
reader_signal: Dict[str, Any],
genre_profile: Dict[str, Any],
low_score_threshold: float,
hook_diversify_enabled: bool,
) -> Dict[str, Any]:
guidance: List[str] = []
low_ranges = reader_signal.get("low_score_ranges") or []
if low_ranges:
worst = min(
low_ranges,
key=lambda row: float(row.get("overall_score", 9999)),
)
guidance.append(
f"第{chapter}章优先修复近期低分段问题:参考{worst.get('start_chapter')}-{worst.get('end_chapter')}章,强化冲突推进与结尾钩子。"
)
hook_usage = reader_signal.get("hook_type_usage") or {}
if hook_usage and hook_diversify_enabled:
dominant_hook = max(hook_usage.items(), key=lambda kv: kv[1])[0]
guidance.append(
f"近期钩子类型“{dominant_hook}”使用偏多,本章建议做钩子差异化,避免连续同构。"
)
pattern_usage = reader_signal.get("pattern_usage") or {}
if pattern_usage:
top_pattern = max(pattern_usage.items(), key=lambda kv: kv[1])[0]
guidance.append(
f"爽点模式“{top_pattern}”近期高频,本章可保留主爽点但叠加一个新爽点副轴。"
)
review_trend = reader_signal.get("review_trend") or {}
overall_avg = review_trend.get("overall_avg")
if isinstance(overall_avg, (int, float)) and float(overall_avg) < low_score_threshold:
guidance.append(
f"最近审查均分{overall_avg:.1f}低于阈值{low_score_threshold:.1f},建议先保稳:减少跳场、每段补动作结果闭环。"
)
genre = str(genre_profile.get("genre") or "").strip()
refs = genre_profile.get("reference_hints") or []
if genre:
guidance.append(f"题材锚定:按“{genre}”叙事主线推进,保持题材读者预期稳定兑现。")
if refs:
guidance.append(f"题材策略可执行提示:{refs[0]}")
guidance.append("网文节奏基线:章首300字内给出目标与阻力,章末保留未闭合问题。")
guidance.append("兑现密度基线:每600-900字给一次微兑现,并确保本章至少1处可量化变化。")
normalized_genre = to_profile_key(genre)
genre_hint = GENRE_GUIDANCE_TEXT.get(normalized_genre)
if genre_hint:
guidance.append(genre_hint)
composite_hints = genre_profile.get("composite_hints") or []
if composite_hints:
guidance.append(f"复合题材协同:{composite_hints[0]}")
if not guidance:
guidance.append("本章执行默认高可读策略:冲突前置、信息后置、段末留钩。")
return {
"guidance": guidance,
"low_ranges": low_ranges,
"hook_usage": hook_usage,
"pattern_usage": pattern_usage,
"genre": genre,
}
def build_writing_checklist(
*,
guidance_items: List[str],
reader_signal: Dict[str, Any],
genre_profile: Dict[str, Any],
strategy_card: Dict[str, Any] | None = None,
min_items: int,
max_items: int,
default_weight: float,
) -> List[Dict[str, Any]]:
items: List[Dict[str, Any]] = []
def _add_item(
item_id: str,
label: str,
*,
weight: float | None = None,
required: bool = False,
source: str = "writing_guidance",
verify_hint: str = "",
) -> None:
if len(items) >= max_items:
return
if any(row.get("id") == item_id for row in items):
return
item_weight = float(weight if weight is not None else default_weight)
if item_weight <= 0:
item_weight = default_weight
items.append(
{
"id": item_id,
"label": label,
"weight": round(item_weight, 2),
"required": bool(required),
"source": source,
"verify_hint": verify_hint,
}
)
low_ranges = reader_signal.get("low_score_ranges") or []
if low_ranges:
worst = min(low_ranges, key=lambda row: float(row.get("overall_score", 9999)))
span = f"{worst.get('start_chapter')}-{worst.get('end_chapter')}"
_add_item(
"fix_low_score_range",
f"修复低分区间问题(参考第{span}章)",
weight=max(default_weight, 1.4),
required=True,
source="reader_signal.low_score_ranges",
verify_hint="至少完成1处冲突升级,并在段末留下钩子。",
)
hook_usage = reader_signal.get("hook_type_usage") or {}
if hook_usage:
dominant_hook = max(hook_usage.items(), key=lambda kv: kv[1])[0]
_add_item(
"hook_diversification",
f"钩子差异化(避免继续单一“{dominant_hook}”)",
weight=max(default_weight, 1.2),
required=True,
source="reader_signal.hook_type_usage",
verify_hint="结尾钩子类型与近20章主类型至少有一处差异。",
)
pattern_usage = reader_signal.get("pattern_usage") or {}
if pattern_usage:
top_pattern = max(pattern_usage.items(), key=lambda kv: kv[1])[0]
_add_item(
"coolpoint_combo",
f"主爽点+副爽点组合(主爽点:{top_pattern})",
weight=default_weight,
required=False,
source="reader_signal.pattern_usage",
verify_hint="新增至少1个副爽点,并与主爽点形成因果链。",
)
review_trend = reader_signal.get("review_trend") or {}
overall_avg = review_trend.get("overall_avg")
if isinstance(overall_avg, (int, float)):
_add_item(
"readability_loop",
"段落可读性闭环(动作→结果→情绪)",
weight=max(default_weight, 1.1),
required=True,
source="reader_signal.review_trend",
verify_hint="抽查3段,均包含动作结果闭环。",
)
genre = str(genre_profile.get("genre") or "").strip()
if genre:
_add_item(
"genre_anchor_consistency",
f"题材锚定一致性({genre})",
weight=max(default_weight, 1.1),
required=True,
source="genre_profile.genre",
verify_hint="主冲突与题材核心承诺保持一致。",
)
if isinstance(strategy_card, dict) and strategy_card.get("enabled"):
_add_item(
"methodology_next_reason",
"方法论:下章动机需可复述(章末或后段均可)",
weight=default_weight,
required=False,
source="methodology.next_reason",
verify_hint="提炼一句“为什么要点下一章”的动机句。",
)
_add_item(
"methodology_power_guard",
"方法论:越级与破局给出机制理由与代价",
weight=default_weight,
required=False,
source="methodology.power_guard",
verify_hint="至少写清1个机制理由与1个代价。"
)
_add_item(
"methodology_antagonist_pressure",
"方法论:反派行动具备目标-手段-代价",
weight=default_weight,
required=False,
source="methodology.antagonist",
verify_hint="反派不是工具人推进,需有可解释行动逻辑。",
)
for idx, text in enumerate(guidance_items, start=1):
if len(items) >= max_items:
break
label = str(text).strip()
if not label:
continue
_add_item(
f"guidance_item_{idx}",
label,
weight=default_weight,
required=False,
source="writing_guidance.guidance_items",
verify_hint="完成后可在正文中定位对应段落。",
)
fallback_items = [
(
"opening_conflict",
"开篇300字内给出冲突触发",
"开头段出现明确目标与阻力。",
),
(
"scene_goal_block",
"场景目标与阻力清晰",
"每个场景至少有1个可验证目标。",
),
(
"ending_hook",
"段末留钩并引出下一问",
"结尾出现未解问题或下一步行动。",
),
]
for item_id, label, verify_hint in fallback_items:
if len(items) >= min_items or len(items) >= max_items:
break
_add_item(
item_id,
label,
weight=default_weight,
required=False,
source="fallback",
verify_hint=verify_hint,
)
return items[:max_items]
def is_checklist_item_completed(item: Dict[str, Any], reader_signal: Dict[str, Any]) -> bool:
item_id = str(item.get("id") or "")
if item_id in {"fix_low_score_range", "readability_loop"}:
review_trend = reader_signal.get("review_trend") or {}
overall = review_trend.get("overall_avg")
return isinstance(overall, (int, float)) and float(overall) >= 75.0
if item_id == "hook_diversification":
hook_usage = reader_signal.get("hook_type_usage") or {}
return len(hook_usage) >= 2
if item_id == "coolpoint_combo":
pattern_usage = reader_signal.get("pattern_usage") or {}
return len(pattern_usage) >= 2
if item_id == "genre_anchor_consistency":
return True
source = str(item.get("source") or "")
if source.startswith("fallback"):
return True
if source.startswith("methodology."):
# 方法论条目当前作为软提示,仅做观察与引导,不参与扣分。
return True
return False