feat: initial commit
This commit is contained in:
commit
ea8f6066c4
217 files changed
+60754
No files matched your search
@@ -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())))
|
||||
|
||||
@@ -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()
|
||||
@@ -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)
|
||||
@@ -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))
|
||||
@@ -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()
|
||||
@@ -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()
|
||||
@@ -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()
|
||||
@@ -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)
|
||||
@@ -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
@@ -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()
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
@@ -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
|
||||
Reference in new issue
Block a user