fix(storage): 严格解析本地策略并原子执行数据库迁移
This commit is contained in:
@@ -32,8 +32,12 @@ def connect() -> sqlite3.Connection:
|
||||
# 关闭 Python sqlite3 的隐式事务,提交时机由 transaction() 或显式 commit 控制。
|
||||
conn.isolation_level = None
|
||||
conn.execute("PRAGMA foreign_keys = ON")
|
||||
_load_extension(conn)
|
||||
migrate(conn)
|
||||
try:
|
||||
_load_extension(conn)
|
||||
migrate(conn)
|
||||
except BaseException:
|
||||
conn.close()
|
||||
raise
|
||||
return conn
|
||||
|
||||
|
||||
|
||||
@@ -6,6 +6,7 @@
|
||||
"""
|
||||
|
||||
from datetime import datetime, timezone
|
||||
import sqlite3
|
||||
|
||||
from app.constants import EMBEDDING_DIM
|
||||
|
||||
@@ -134,6 +135,18 @@ MIGRATIONS: list[str] = [
|
||||
]
|
||||
|
||||
|
||||
def _statements(script: str):
|
||||
"""Split complete SQLite statements without executescript's implicit COMMIT."""
|
||||
pending = ""
|
||||
for char in script:
|
||||
pending += char
|
||||
if char == ";" and sqlite3.complete_statement(pending):
|
||||
yield pending
|
||||
pending = ""
|
||||
if pending.strip():
|
||||
yield pending
|
||||
|
||||
|
||||
def migrate(conn) -> None:
|
||||
"""把尚未应用的迁移脚本按序应用到给定连接。"""
|
||||
conn.execute(
|
||||
@@ -145,9 +158,28 @@ def migrate(conn) -> None:
|
||||
for idx, script in enumerate(MIGRATIONS, start=1):
|
||||
if idx in applied:
|
||||
continue
|
||||
conn.executescript(script)
|
||||
conn.execute(
|
||||
"INSERT INTO schema_migrations (version, applied_at) VALUES (?, ?)",
|
||||
(idx, datetime.now(timezone.utc).isoformat()),
|
||||
)
|
||||
conn.commit()
|
||||
conn.execute("BEGIN IMMEDIATE")
|
||||
try:
|
||||
# Another connection may have migrated while this one waited.
|
||||
if not conn.execute("SELECT 1 FROM schema_migrations WHERE version=?", (idx,)).fetchone():
|
||||
recovered_v6 = False
|
||||
if idx == 6:
|
||||
column = next((row for row in conn.execute("PRAGMA table_info(blocks)")
|
||||
if row["name"] == "embedding_local_only"), None)
|
||||
if column is not None:
|
||||
# Recover the precise partial state left by the old v6 runner.
|
||||
if column["type"].upper() != "INTEGER" or column["notnull"] != 1 or column["dflt_value"] != "0":
|
||||
raise sqlite3.DatabaseError("Unexpected embedding_local_only column schema")
|
||||
recovered_v6 = True
|
||||
if not recovered_v6:
|
||||
for statement in _statements(script):
|
||||
conn.execute(statement)
|
||||
conn.execute(
|
||||
"INSERT INTO schema_migrations (version, applied_at) VALUES (?, ?)",
|
||||
(idx, datetime.now(timezone.utc).isoformat()),
|
||||
)
|
||||
conn.execute("COMMIT")
|
||||
except BaseException:
|
||||
if conn.in_transaction:
|
||||
conn.execute("ROLLBACK")
|
||||
raise
|
||||
|
||||
@@ -13,7 +13,10 @@ from dataclasses import dataclass, field
|
||||
from datetime import datetime
|
||||
from pathlib import Path
|
||||
|
||||
import yaml
|
||||
|
||||
from app.contracts import NoteBlock
|
||||
from app.errors import ApiError
|
||||
from app.textutils import count_tokens
|
||||
|
||||
_HEADING_RE = re.compile(r"^(#{1,6})[ \t]+(.*?)\s*$")
|
||||
@@ -70,7 +73,7 @@ def parse_note(
|
||||
created_at=created_at,
|
||||
updated_at=updated_at,
|
||||
blocks=blocks,
|
||||
embedding_local_only=str(frontmatter.get("embedding_local_only", "")).lower() == "true",
|
||||
embedding_local_only=_embedding_policy(markdown),
|
||||
)
|
||||
|
||||
|
||||
@@ -184,6 +187,35 @@ def _utf16_len(text: str) -> int:
|
||||
return len(text.encode("utf-16-le")) // 2
|
||||
|
||||
|
||||
def _embedding_policy(markdown: str) -> bool:
|
||||
end = markdown.find("\n---", 3) if markdown.startswith("---") else -1
|
||||
if end == -1:
|
||||
return False
|
||||
try:
|
||||
# Compose nodes without constructing objects. This accepts YAML comments,
|
||||
# quoted keys and indentation while retaining duplicate-key information.
|
||||
node = yaml.compose(markdown[3:end], Loader=yaml.SafeLoader)
|
||||
except yaml.YAMLError as exc:
|
||||
raise ApiError(422, "INVALID_EMBEDDING_POLICY", "Frontmatter YAML 无效,无法确认本地索引策略。") from exc
|
||||
if node is None:
|
||||
return False
|
||||
if not isinstance(node, yaml.MappingNode):
|
||||
raise ApiError(422, "INVALID_EMBEDDING_POLICY", "Frontmatter 必须是 YAML 键值映射。")
|
||||
if any(key.tag == "tag:yaml.org,2002:merge" for key, _ in node.value):
|
||||
raise ApiError(422, "INVALID_EMBEDDING_POLICY", "Frontmatter 不支持 YAML 合并键,请显式声明索引策略。")
|
||||
values = [value for key, value in node.value
|
||||
if isinstance(key, yaml.ScalarNode) and key.value.lower() == "embedding_local_only"]
|
||||
if not values:
|
||||
return False
|
||||
if len(values) > 1:
|
||||
raise ApiError(422, "INVALID_EMBEDDING_POLICY", "embedding_local_only 不能重复声明。")
|
||||
value = values[0]
|
||||
if (not isinstance(value, yaml.ScalarNode) or value.tag != "tag:yaml.org,2002:bool"
|
||||
or value.value.lower() not in {"true", "false", "yes", "no", "on", "off"}):
|
||||
raise ApiError(422, "INVALID_EMBEDDING_POLICY", "embedding_local_only 必须是 YAML 布尔值 true 或 false。")
|
||||
return value.value.lower() in {"true", "yes", "on"}
|
||||
|
||||
|
||||
def _extract_frontmatter(markdown: str) -> dict[str, str]:
|
||||
"""极简 frontmatter 解析,只提取 key: value 行。"""
|
||||
if not markdown.startswith("---"):
|
||||
|
||||
Reference in New Issue
Block a user