Files
daily_stock_analysis/tests/test_storage.py
Elvis Wang 1b42907605 feat: 扩展 A 股指数注册表与 CSI provider 路由 (#2267)
* feat: add JSON-backed index registry

* fix: address index registry review feedback

* fix: validate raw duplicate index aliases

* fix: converge CSI task identities and registry fallback

* docs: align index identity terminology
2026-08-25 22:32:04 +08:00

2067 lines
87 KiB
Python
Raw Permalink Blame History

This file contains ambiguous Unicode characters

This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.

# -*- coding: utf-8 -*-
import unittest
import sys
import os
import shutil
import sqlite3
import tempfile
import threading
from datetime import date
from unittest.mock import patch
import pandas as pd
from sqlalchemy import and_, create_engine as sqlalchemy_create_engine, inspect, select
from sqlalchemy.sql import func
# Ensure src module can be imported
sys.path.insert(0, os.path.abspath(os.path.join(os.path.dirname(__file__), '..')))
from src.config import Config
from src.storage import Base, CURRENT_SCHEMA_VERSION, DatabaseManager, DatabaseSchemaMigration, StockDaily
from src.services import stock_list_parser as stock_list_parser_module
class TestStorage(unittest.TestCase):
@staticmethod
def _list_sqlite_indexes(db_path: str, table_name: str) -> dict[str, list[str]]:
with sqlite3.connect(db_path) as conn:
indexes = {}
for row in conn.execute(f"PRAGMA index_list({table_name})").fetchall():
index_name = row[1]
indexes[index_name] = [
index_info[2]
for index_info in conn.execute(
f"PRAGMA index_xinfo({index_name})"
).fetchall()
if index_info[2] is not None and int(index_info[5]) == 1
]
return indexes
@staticmethod
def _list_sqlite_unique_indexes(db_path: str, table_name: str) -> dict[str, list[str]]:
with sqlite3.connect(db_path) as conn:
rows = conn.execute(f"PRAGMA index_list({table_name})").fetchall()
unique_indexes = {}
for row in rows:
if int(row[2]) != 1:
continue
index_name = row[1]
index_columns = []
for index_info in conn.execute(f"PRAGMA index_xinfo({index_name})").fetchall():
column_name = index_info[2]
if column_name is not None:
index_columns.append(column_name)
unique_indexes[index_name] = index_columns
return unique_indexes
def test_legacy_intelligence_items_url_unique_index_rebuilds_without_collision(self) -> None:
temp_dir = tempfile.TemporaryDirectory()
db_path = os.path.join(temp_dir.name, "legacy_intel.sqlite")
try:
with sqlite3.connect(db_path) as conn:
conn.execute(
"""CREATE TABLE intelligence_sources (
id INTEGER PRIMARY KEY AUTOINCREMENT,
name TEXT NOT NULL UNIQUE,
source_type TEXT NOT NULL DEFAULT 'rss',
url TEXT NOT NULL,
scope_type TEXT NOT NULL DEFAULT 'market',
scope_value TEXT,
market TEXT NOT NULL DEFAULT 'cn',
enabled INTEGER NOT NULL DEFAULT 1,
last_status TEXT,
last_error TEXT,
last_fetched_at DATETIME,
created_at DATETIME,
updated_at DATETIME
)"""
)
conn.execute(
"""CREATE TABLE intelligence_items (
id INTEGER PRIMARY KEY AUTOINCREMENT,
source_id INTEGER,
source_name TEXT,
source_type TEXT NOT NULL DEFAULT 'rss',
title TEXT NOT NULL,
summary TEXT,
url TEXT NOT NULL,
source TEXT,
published_at DATETIME,
fetched_at DATETIME,
scope_type TEXT NOT NULL DEFAULT 'market',
scope_value TEXT,
market TEXT NOT NULL DEFAULT 'cn',
raw_payload TEXT
)"""
)
conn.execute("CREATE UNIQUE INDEX uix_intelligence_item_url_legacy ON intelligence_items(url)")
conn.execute("CREATE INDEX ix_intel_item_scope_time ON intelligence_items(scope_type, scope_value, market, published_at)")
conn.execute("CREATE INDEX ix_intel_item_fetch_time ON intelligence_items(fetched_at)")
conn.execute("INSERT INTO intelligence_sources (name, url) VALUES ('legacy', 'https://legacy.example.com/rss.xml')")
source_id = conn.execute("SELECT id FROM intelligence_sources WHERE name='legacy'").fetchone()[0]
conn.executemany(
"""INSERT INTO intelligence_items (
source_id,
source_name,
source_type,
title,
summary,
url,
source,
published_at,
fetched_at,
scope_type,
scope_value,
market,
raw_payload
) VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)""",
[
(source_id, 'legacy-source', 'rss', 'A', 'legacy-a', 'https://legacy.example.com/a', 'legacy',
'2026-01-01 00:00:00', '2026-01-01 00:00:00', 'market', None, 'cn', None),
(source_id, 'legacy-source', 'rss', 'B', 'legacy-b', 'https://legacy.example.com/b', 'legacy',
'2026-01-02 00:00:00', '2026-01-02 00:00:00', 'market', None, 'cn', None),
],
)
unique_indexes_before = self._list_sqlite_unique_indexes(db_path, "intelligence_items")
self.assertIn("uix_intelligence_item_url_legacy", unique_indexes_before)
self.assertEqual(unique_indexes_before["uix_intelligence_item_url_legacy"], ["url"])
DatabaseManager.reset_instance()
Config.reset_instance()
DatabaseManager(db_url=f"sqlite:///{db_path}")
unique_indexes_after = self._list_sqlite_unique_indexes(db_path, "intelligence_items")
self.assertNotIn("uix_intelligence_item_url_legacy", unique_indexes_after)
self.assertIn("uix_intel_item_scope", unique_indexes_after)
self.assertEqual(
unique_indexes_after["uix_intel_item_scope"],
["source_id", "url", "scope_type", "scope_value", "market"],
)
with sqlite3.connect(db_path) as conn:
table_count = conn.execute("SELECT COUNT(*) FROM intelligence_items").fetchone()[0]
temp_tables = conn.execute(
"SELECT name FROM sqlite_master WHERE type='table' AND name LIKE 'intelligence_items_recreate_tmp_%'"
).fetchall()
self.assertEqual(table_count, 2)
self.assertEqual(temp_tables, [])
finally:
DatabaseManager.reset_instance()
Config.reset_instance()
temp_dir.cleanup()
def test_database_initialization_records_schema_version(self):
DatabaseManager.reset_instance()
db = DatabaseManager(db_url="sqlite:///:memory:")
with db.get_session() as session:
row = session.get(DatabaseSchemaMigration, CURRENT_SCHEMA_VERSION)
self.assertIsNotNone(row)
self.assertEqual(row.version, CURRENT_SCHEMA_VERSION)
self.assertIn("metadata.create_all", row.description)
DatabaseManager.reset_instance()
def test_schema_migration_record_is_idempotent(self):
DatabaseManager.reset_instance()
db = DatabaseManager(db_url="sqlite:///:memory:")
db._ensure_schema_migration_record()
db._ensure_schema_migration_record()
with db.get_session() as session:
count = session.execute(
select(func.count()).select_from(DatabaseSchemaMigration)
).scalar_one()
self.assertEqual(count, 1)
DatabaseManager.reset_instance()
def test_fresh_decision_signal_schema_has_profile_indexes(self):
DatabaseManager.reset_instance()
temp_dir = tempfile.TemporaryDirectory()
db_path = os.path.join(temp_dir.name, "fresh_decision_profile.db")
try:
DatabaseManager(db_url=f"sqlite:///{db_path}")
indexes = self._list_sqlite_indexes(db_path, "decision_signals")
self.assertEqual(
indexes.get("ix_decision_signals_decision_profile"),
["decision_profile"],
)
self.assertEqual(
indexes.get("ix_decision_signal_market_stock_profile_created"),
["market", "stock_code", "decision_profile", "created_at"],
)
self.assertEqual(
indexes.get(
"ix_decision_signal_report_type_market_stock_profile_action_horizon_phase"
),
[
"source_report_id", "source_type", "market", "stock_code",
"decision_profile", "action", "horizon", "market_phase",
],
)
self.assertEqual(
indexes.get(
"ix_decision_signal_trace_type_market_stock_profile_action_horizon_phase"
),
[
"trace_id", "source_type", "market", "stock_code",
"decision_profile", "action", "horizon", "market_phase",
],
)
finally:
DatabaseManager.reset_instance()
Config.reset_instance()
temp_dir.cleanup()
def test_decision_signal_profile_migration_adds_column_indexes_and_closed_stats(self):
DatabaseManager.reset_instance()
temp_dir = tempfile.TemporaryDirectory()
db_path = os.path.join(temp_dir.name, "legacy_decision_profile.db")
deeply_nested_json = "[" * 10_000 + "]" * 10_000
try:
with sqlite3.connect(db_path) as conn:
conn.execute(
"""CREATE TABLE decision_signals (
id INTEGER PRIMARY KEY AUTOINCREMENT,
stock_code TEXT,
market TEXT,
source_type TEXT,
source_report_id INTEGER,
trace_id TEXT,
action TEXT,
horizon TEXT,
market_phase TEXT,
created_at DATETIME,
metadata_json TEXT
)"""
)
conn.execute(
"CREATE INDEX ix_decision_signal_report_type_market_stock_action_horizon_phase "
"ON decision_signals "
"(source_report_id, source_type, market, stock_code, action, horizon, market_phase)"
)
conn.execute(
"CREATE INDEX ix_decision_signal_trace_type_market_stock_action_horizon_phase "
"ON decision_signals "
"(trace_id, source_type, market, stock_code, action, horizon, market_phase)"
)
conn.executemany(
"""INSERT INTO decision_signals (
stock_code,
market,
source_type,
source_report_id,
trace_id,
action,
horizon,
market_phase,
created_at,
metadata_json
) VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?)""",
[
("600519", "cn", "analysis", 1, "trace-1", "buy", "3d", "intraday", "2026-01-01", '{"decision_profile":"balanced"}'),
("600519", "cn", "analysis", 2, "trace-2", "buy", "3d", "intraday", "2026-01-01", None),
("600519", "cn", "analysis", 3, "trace-3", "buy", "3d", "intraday", "2026-01-01", ""),
("600519", "cn", "analysis", 4, "trace-4", "buy", "3d", "intraday", "2026-01-01", " "),
("600519", "cn", "analysis", 5, "trace-5", "buy", "3d", "intraday", "2026-01-01", "{not-json"),
("600519", "cn", "analysis", 6, "trace-6", "buy", "3d", "intraday", "2026-01-01", sqlite3.Binary(b"\xff")),
("600519", "cn", "analysis", 7, "trace-7", "buy", "3d", "intraday", "2026-01-01", "null"),
("600519", "cn", "analysis", 8, "trace-8", "buy", "3d", "intraday", "2026-01-01", "[]"),
("600519", "cn", "analysis", 9, "trace-9", "buy", "3d", "intraday", "2026-01-01", '"balanced"'),
("600519", "cn", "analysis", 10, "trace-10", "buy", "3d", "intraday", "2026-01-01", "1"),
("600519", "cn", "analysis", 11, "trace-11", "buy", "3d", "intraday", "2026-01-01", "{}"),
("600519", "cn", "analysis", 12, "trace-12", "buy", "3d", "intraday", "2026-01-01", '{"decision_profile":null}'),
("600519", "cn", "analysis", 13, "trace-13", "buy", "3d", "intraday", "2026-01-01", '{"decision_profile":""}'),
("600519", "cn", "analysis", 14, "trace-14", "buy", "3d", "intraday", "2026-01-01", '{"decision_profile":" "}'),
("600519", "cn", "analysis", 15, "trace-15", "buy", "3d", "intraday", "2026-01-01", '{"decision_profile":"reckless"}'),
("600519", "cn", "analysis", 16, "trace-16", "buy", "3d", "intraday", "2026-01-01", deeply_nested_json),
],
)
with self.assertLogs("src.storage", level="INFO") as logs:
DatabaseManager(db_url=f"sqlite:///{db_path}")
with sqlite3.connect(db_path) as conn:
columns = {row[1] for row in conn.execute("PRAGMA table_info(decision_signals)").fetchall()}
rows = conn.execute(
"SELECT id, decision_profile FROM decision_signals ORDER BY id"
).fetchall()
self.assertIn("decision_profile", columns)
self.assertEqual(rows[0], (1, "balanced"))
self.assertTrue(all(profile is None for _, profile in rows[1:]))
indexes = self._list_sqlite_indexes(db_path, "decision_signals")
expected_indexes = {
"ix_decision_signals_decision_profile": ["decision_profile"],
"ix_decision_signal_market_stock_profile_created": [
"market", "stock_code", "decision_profile", "created_at",
],
"ix_decision_signal_report_type_market_stock_profile_action_horizon_phase": [
"source_report_id", "source_type", "market", "stock_code",
"decision_profile", "action", "horizon", "market_phase",
],
"ix_decision_signal_trace_type_market_stock_profile_action_horizon_phase": [
"trace_id", "source_type", "market", "stock_code",
"decision_profile", "action", "horizon", "market_phase",
],
}
for index_name, index_columns in expected_indexes.items():
self.assertEqual(indexes.get(index_name), index_columns)
self.assertEqual(
indexes.get("ix_decision_signal_report_type_market_stock_action_horizon_phase"),
[
"source_report_id", "source_type", "market", "stock_code",
"action", "horizon", "market_phase",
],
)
self.assertEqual(
indexes.get("ix_decision_signal_trace_type_market_stock_action_horizon_phase"),
[
"trace_id", "source_type", "market", "stock_code",
"action", "horizon", "market_phase",
],
)
log_text = "\n".join(logs.output)
self.assertIn("candidate_count=16", log_text)
self.assertIn("backfilled_count=1", log_text)
self.assertIn("guard_skipped_count=0", log_text)
self.assertIn("missing_metadata_count=1", log_text)
self.assertIn("missing_profile_count=4", log_text)
self.assertIn("invalid_json_count=5", log_text)
self.assertIn("non_object_count=4", log_text)
self.assertIn("invalid_profile_count=1", log_text)
self.assertIn("skipped_existing_profile_count=0", log_text)
DatabaseManager.reset_instance()
with self.assertLogs("src.storage", level="INFO") as second_logs:
DatabaseManager(db_url=f"sqlite:///{db_path}")
second_log_text = "\n".join(second_logs.output)
self.assertIn("candidate_count=15", second_log_text)
self.assertIn("backfilled_count=0", second_log_text)
self.assertIn("guard_skipped_count=0", second_log_text)
self.assertIn("missing_metadata_count=1", second_log_text)
self.assertIn("missing_profile_count=4", second_log_text)
self.assertIn("invalid_json_count=5", second_log_text)
self.assertIn("non_object_count=4", second_log_text)
self.assertIn("invalid_profile_count=1", second_log_text)
self.assertIn("skipped_existing_profile_count=1", second_log_text)
finally:
DatabaseManager.reset_instance()
Config.reset_instance()
temp_dir.cleanup()
def test_decision_signal_profile_migration_runs_when_column_already_exists(self):
DatabaseManager.reset_instance()
temp_dir = tempfile.TemporaryDirectory()
db_path = os.path.join(temp_dir.name, "existing_decision_profile.db")
try:
with sqlite3.connect(db_path) as conn:
conn.execute(
"""CREATE TABLE decision_signals (
id INTEGER PRIMARY KEY AUTOINCREMENT,
stock_code TEXT,
market TEXT,
source_type TEXT,
source_report_id INTEGER,
trace_id TEXT,
action TEXT,
horizon TEXT,
market_phase TEXT,
created_at DATETIME,
metadata_json TEXT,
decision_profile VARCHAR(16)
)"""
)
conn.executemany(
"INSERT INTO decision_signals (metadata_json, decision_profile) VALUES (?, ?)",
[
('{"decision_profile":"aggressive"}', None),
(None, None),
('{"decision_profile":"balanced"}', "conservative"),
('{"decision_profile":"balanced"}', ""),
],
)
with self.assertLogs("src.storage", level="INFO") as logs:
DatabaseManager(db_url=f"sqlite:///{db_path}")
with sqlite3.connect(db_path) as conn:
profiles = conn.execute(
"SELECT decision_profile FROM decision_signals ORDER BY id"
).fetchall()
self.assertEqual(
profiles,
[("aggressive",), (None,), ("conservative",), ("",)],
)
log_text = "\n".join(logs.output)
self.assertIn("candidate_count=2", log_text)
self.assertIn("backfilled_count=1", log_text)
self.assertIn("missing_metadata_count=1", log_text)
self.assertIn("skipped_existing_profile_count=2", log_text)
finally:
DatabaseManager.reset_instance()
Config.reset_instance()
temp_dir.cleanup()
def test_decision_signal_profile_migration_fails_when_column_inspection_fails(self):
class BrokenInspector:
def has_table(self, _table_name: str) -> bool:
return True
def get_columns(self, _table_name: str):
raise RuntimeError("inspection failed")
DatabaseManager.reset_instance()
try:
with patch("src.storage.inspect", return_value=BrokenInspector()):
with self.assertLogs("src.storage", level="ERROR") as logs:
with self.assertRaises(RuntimeError):
DatabaseManager(db_url="sqlite:///:memory:")
self.assertIn(
"profile migration cannot continue safely",
"\n".join(logs.output),
)
finally:
DatabaseManager.reset_instance()
Config.reset_instance()
def test_schema_migration_record_handles_concurrent_initialization(self):
DatabaseManager.reset_instance()
temp_dir = tempfile.TemporaryDirectory()
db_path = os.path.join(temp_dir.name, "schema_migration_race.db")
db = DatabaseManager(db_url=f"sqlite:///{db_path}")
worker_count = 8
barrier = threading.Barrier(worker_count)
errors = []
state_lock = threading.Lock()
with db.get_session() as session:
session.query(DatabaseSchemaMigration).delete()
session.commit()
def ensure_record() -> None:
try:
barrier.wait(timeout=5)
db._ensure_schema_migration_record()
except Exception as exc:
with state_lock:
errors.append(exc)
threads = [threading.Thread(target=ensure_record) for _ in range(worker_count)]
for thread in threads:
thread.start()
for thread in threads:
thread.join(timeout=5)
with db.get_session() as session:
rows = session.execute(select(DatabaseSchemaMigration)).scalars().all()
self.assertFalse(any(thread.is_alive() for thread in threads))
self.assertEqual(errors, [])
self.assertEqual(len(rows), 1)
self.assertEqual(rows[0].version, CURRENT_SCHEMA_VERSION)
DatabaseManager.reset_instance()
temp_dir.cleanup()
def test_parse_sniper_value(self):
"""测试解析狙击点位数值"""
# 1. 正常数值
self.assertEqual(DatabaseManager._parse_sniper_value(100), 100.0)
self.assertEqual(DatabaseManager._parse_sniper_value(100.5), 100.5)
self.assertEqual(DatabaseManager._parse_sniper_value("100"), 100.0)
self.assertEqual(DatabaseManager._parse_sniper_value("100.5"), 100.5)
# 2. 包含中文描述和"元"
self.assertEqual(DatabaseManager._parse_sniper_value("建议在 100 元附近买入"), 100.0)
self.assertEqual(DatabaseManager._parse_sniper_value("价格100.5元"), 100.5)
# 3. 包含干扰数字修复的Bug场景
# 之前 "MA5" 会被错误提取为 5.0,现在应该提取 "元" 前面的 100
text_bug = "无法给出。需等待MA5数据恢复在股价回踩MA5且乖离率<2%时考虑100元"
self.assertEqual(DatabaseManager._parse_sniper_value(text_bug), 100.0)
# 4. 更多干扰场景
text_complex = "MA10为20.5建议在30元买入"
self.assertEqual(DatabaseManager._parse_sniper_value(text_complex), 30.0)
text_multiple = "支撑位10元阻力位20元" # 应该提取最后一个"元"前面的数字即20或者更复杂的逻辑
# 当前逻辑是找最后一个冒号,然后找之后的第一个"元",提取中间的数字。
# 测试没有冒号的情况
self.assertEqual(DatabaseManager._parse_sniper_value("30元"), 30.0)
# 测试多个数字在"元"之前
self.assertEqual(DatabaseManager._parse_sniper_value("MA5 10 20元"), 20.0)
# 5. Fallback: no "元" character — extracts last non-MA number
self.assertEqual(DatabaseManager._parse_sniper_value("102.10-103.00MA5附近"), 103.0)
self.assertEqual(DatabaseManager._parse_sniper_value("97.62-98.50MA10附近"), 98.5)
self.assertEqual(DatabaseManager._parse_sniper_value("93.40下方MA20支撑"), 93.4)
self.assertEqual(DatabaseManager._parse_sniper_value("108.00-110.00(前期高点阻力)"), 110.0)
# 6. 无效输入
self.assertIsNone(DatabaseManager._parse_sniper_value(None))
self.assertIsNone(DatabaseManager._parse_sniper_value(""))
self.assertIsNone(DatabaseManager._parse_sniper_value("没有数字"))
self.assertIsNone(DatabaseManager._parse_sniper_value("MA5但没有元"))
# 7. 回归:括号内技术指标数字不应被提取
self.assertNotEqual(DatabaseManager._parse_sniper_value("1.52-1.53 (回踩MA5/10附近)"), 10.0)
self.assertNotEqual(DatabaseManager._parse_sniper_value("1.55-1.56(MA5/M20支撑)"), 20.0)
self.assertNotEqual(DatabaseManager._parse_sniper_value("1.49-1.50(MA60附近企稳)"), 60.0)
# 验证正确值在区间内
self.assertIn(DatabaseManager._parse_sniper_value("1.52-1.53 (回踩MA5/10附近)"), [1.52, 1.53])
self.assertIn(DatabaseManager._parse_sniper_value("1.55-1.56(MA5/M20支撑)"), [1.55, 1.56])
self.assertIn(DatabaseManager._parse_sniper_value("1.49-1.50(MA60附近企稳)"), [1.49, 1.50])
def test_get_chat_sessions_prefix_is_scoped_by_colon_boundary(self):
DatabaseManager.reset_instance()
db = DatabaseManager(db_url="sqlite:///:memory:")
db.save_conversation_message("telegram_12345:chat", "user", "first user")
db.save_conversation_message("telegram_123456:chat", "user", "second user")
sessions = db.get_chat_sessions(session_prefix="telegram_12345")
self.assertEqual(len(sessions), 1)
self.assertEqual(sessions[0]["session_id"], "telegram_12345:chat")
DatabaseManager.reset_instance()
def test_get_chat_sessions_can_include_legacy_exact_session_id(self):
DatabaseManager.reset_instance()
db = DatabaseManager(db_url="sqlite:///:memory:")
db.save_conversation_message("feishu_u1", "user", "legacy chat")
db.save_conversation_message("feishu_u1:ask_600519", "user", "ask session")
sessions = db.get_chat_sessions(
session_prefix="feishu_u1:",
extra_session_ids=["feishu_u1"],
)
self.assertEqual({item["session_id"] for item in sessions}, {"feishu_u1", "feishu_u1:ask_600519"})
DatabaseManager.reset_instance()
def test_conversation_summary_upsert_and_delete_with_session(self):
DatabaseManager.reset_instance()
db = DatabaseManager(db_url="sqlite:///:memory:")
db.save_conversation_message("summary-session", "user", "hello")
db.upsert_conversation_summary(
"summary-session",
"first summary",
covered_message_id=1,
source_message_count=1,
estimated_tokens=10,
)
db.upsert_conversation_summary(
"summary-session",
"updated summary",
covered_message_id=2,
source_message_count=2,
estimated_tokens=12,
)
summary = db.get_conversation_summary("summary-session")
self.assertIsNotNone(summary)
self.assertEqual(summary["summary"], "updated summary")
self.assertEqual(summary["covered_message_id"], 2)
self.assertEqual(summary["source_message_count"], 2)
deleted = db.delete_conversation_session("summary-session")
self.assertEqual(deleted, 1)
self.assertIsNone(db.get_conversation_summary("summary-session"))
DatabaseManager.reset_instance()
def test_conversation_message_save_returns_id(self):
DatabaseManager.reset_instance()
db = DatabaseManager(db_url="sqlite:///:memory:")
message_id = db.save_conversation_message("message-id-session", "user", "hello")
self.assertIsInstance(message_id, int)
self.assertGreater(message_id, 0)
DatabaseManager.reset_instance()
def test_conversation_user_turn_persists_and_updates_session_skills(self):
DatabaseManager.reset_instance()
db = DatabaseManager(db_url="sqlite:///:memory:")
first_id = db.save_conversation_user_turn(
"skill-session",
"first question",
["technical", "risk"],
)
second_id = db.save_conversation_user_turn(
"skill-session",
"use general analysis",
[],
)
self.assertGreater(first_id, 0)
self.assertGreater(second_id, first_id)
self.assertEqual(
[message["content"] for message in db.get_conversation_messages("skill-session")],
["first question", "use general analysis"],
)
self.assertEqual(
db.get_conversation_session_selected_skill_ids("skill-session"),
[],
)
deleted = db.delete_conversation_session("skill-session")
self.assertEqual(deleted, 2)
self.assertIsNone(
db.get_conversation_session_selected_skill_ids("skill-session")
)
DatabaseManager.reset_instance()
def test_conversation_user_turn_without_skill_update_keeps_session_state(self):
DatabaseManager.reset_instance()
db = DatabaseManager(db_url="sqlite:///:memory:")
db.save_conversation_user_turn("skill-session", "first", ["technical"])
db.save_conversation_user_turn("skill-session", "follow up")
self.assertEqual(
db.get_conversation_session_selected_skill_ids("skill-session"),
["technical"],
)
DatabaseManager.reset_instance()
def test_conversation_user_turn_rolls_back_message_when_state_write_fails(self):
DatabaseManager.reset_instance()
db = DatabaseManager(db_url="sqlite:///:memory:")
with patch("src.storage.sqlite_insert", side_effect=RuntimeError("state write failed")):
with self.assertRaisesRegex(RuntimeError, "state write failed"):
db.save_conversation_user_turn(
"skill-session",
"not accepted",
["technical"],
)
self.assertEqual(db.get_conversation_messages("skill-session"), [])
self.assertIsNone(
db.get_conversation_session_selected_skill_ids("skill-session")
)
DatabaseManager.reset_instance()
def test_provider_turn_round_trip_preserves_protocol_fields_and_flags(self):
DatabaseManager.reset_instance()
db = DatabaseManager(db_url="sqlite:///:memory:")
user_id = db.save_conversation_message("trace-session", "user", "question")
assistant_id = db.save_conversation_message("trace-session", "assistant", "final")
trace_messages = [
{
"role": "assistant",
"content": "checking",
"reasoning_content": "reasoning",
"tool_calls": [
{
"id": "call_1",
"name": "echo",
"arguments": {"message": "hello"},
"provider_specific_fields": {"thought_signature": "sig"},
}
],
},
{"role": "tool", "tool_call_id": "call_1", "content": "{\"ok\": true}"},
]
turn_id = db.save_agent_provider_turn(
session_id="trace-session",
run_id="run-1",
provider="deepseek",
model="deepseek/deepseek-chat",
anchor_user_message_id=user_id,
anchor_assistant_message_id=assistant_id,
messages=trace_messages,
contains_reasoning=True,
contains_tool_calls=True,
contains_thinking_blocks=False,
must_roundtrip=True,
estimated_tokens=42,
)
rows = db.get_agent_provider_turns("trace-session")
self.assertIsInstance(turn_id, int)
self.assertEqual(len(rows), 1)
self.assertEqual(rows[0]["messages"], trace_messages)
self.assertTrue(rows[0]["contains_reasoning"])
self.assertTrue(rows[0]["contains_tool_calls"])
self.assertTrue(rows[0]["must_roundtrip"])
self.assertEqual(rows[0]["estimated_tokens"], 42)
DatabaseManager.reset_instance()
def test_provider_turns_do_not_appear_in_visible_or_web_messages_and_delete_with_session(self):
DatabaseManager.reset_instance()
db = DatabaseManager(db_url="sqlite:///:memory:")
user_id = db.save_conversation_message("trace-hidden", "user", "visible question")
assistant_id = db.save_conversation_message("trace-hidden", "assistant", "visible answer")
db.save_agent_provider_turn(
session_id="trace-hidden",
run_id="run-hidden",
provider="deepseek",
model="deepseek/deepseek-chat",
anchor_user_message_id=user_id,
anchor_assistant_message_id=assistant_id,
messages=[{"role": "assistant", "reasoning_content": "SECRET_REASONING", "tool_calls": []}],
contains_reasoning=True,
contains_tool_calls=True,
contains_thinking_blocks=False,
must_roundtrip=True,
estimated_tokens=5,
)
self.assertEqual(
[(m["role"], m["content"]) for m in db.get_visible_conversation_messages("trace-hidden")],
[("user", "visible question"), ("assistant", "visible answer")],
)
self.assertEqual(
[(m["role"], m["content"]) for m in db.get_conversation_history("trace-hidden")],
[("user", "visible question"), ("assistant", "visible answer")],
)
self.assertEqual(
[(m["role"], m["content"]) for m in db.get_conversation_messages("trace-hidden")],
[("user", "visible question"), ("assistant", "visible answer")],
)
deleted = db.delete_conversation_session("trace-hidden")
self.assertEqual(deleted, 2)
self.assertEqual(db.get_agent_provider_turns("trace-hidden"), [])
DatabaseManager.reset_instance()
def test_provider_turn_retention_is_bucketed_by_session_provider_model(self):
DatabaseManager.reset_instance()
db = DatabaseManager(db_url="sqlite:///:memory:")
for idx in range(5):
user_id = db.save_conversation_message("retention", "user", f"q{idx}")
assistant_id = db.save_conversation_message("retention", "assistant", f"a{idx}")
db.save_agent_provider_turn(
session_id="retention",
run_id=f"run-{idx}",
provider="deepseek",
model="deepseek/deepseek-chat",
anchor_user_message_id=user_id,
anchor_assistant_message_id=assistant_id,
messages=[{"role": "assistant", "reasoning_content": f"r{idx}", "tool_calls": [{"id": f"c{idx}", "name": "echo", "arguments": {}}]}],
contains_reasoning=True,
contains_tool_calls=True,
contains_thinking_blocks=False,
must_roundtrip=True,
estimated_tokens=idx + 1,
)
user_id = db.save_conversation_message("retention", "user", "other")
assistant_id = db.save_conversation_message("retention", "assistant", "other")
db.save_agent_provider_turn(
session_id="retention",
run_id="run-other",
provider="anthropic",
model="anthropic/claude-test",
anchor_user_message_id=user_id,
anchor_assistant_message_id=assistant_id,
messages=[{"role": "assistant", "provider_blocks": [{"type": "thinking"}], "tool_calls": [{"id": "c-other", "name": "echo", "arguments": {}}]}],
contains_reasoning=False,
contains_tool_calls=True,
contains_thinking_blocks=True,
must_roundtrip=True,
estimated_tokens=1,
)
deepseek_rows = db.get_agent_provider_turns(
"retention",
provider="deepseek",
model="deepseek/deepseek-chat",
)
anthropic_rows = db.get_agent_provider_turns(
"retention",
provider="anthropic",
model="anthropic/claude-test",
)
self.assertEqual(len(deepseek_rows), 3)
self.assertEqual([row["run_id"] for row in deepseek_rows], ["run-2", "run-3", "run-4"])
self.assertEqual(len(anthropic_rows), 1)
DatabaseManager.reset_instance()
def test_get_visible_conversation_messages_returns_ordered_visible_content(self):
DatabaseManager.reset_instance()
db = DatabaseManager(db_url="sqlite:///:memory:")
db.save_conversation_message("visible-session", "system", "hidden")
db.save_conversation_message("visible-session", "user", "question")
db.save_conversation_message("visible-session", "assistant", "answer")
messages = db.get_visible_conversation_messages("visible-session")
self.assertEqual(
[(item["role"], item["content"]) for item in messages],
[("user", "question"), ("assistant", "answer")],
)
self.assertIsInstance(messages[0]["id"], int)
DatabaseManager.reset_instance()
def test_get_visible_conversation_messages_limit_returns_ordered_tail(self):
DatabaseManager.reset_instance()
db = DatabaseManager(db_url="sqlite:///:memory:")
for idx in range(25):
db.save_conversation_message("visible-limit", "user", f"msg-{idx}")
messages = db.get_visible_conversation_messages("visible-limit", limit=20)
self.assertEqual(len(messages), 20)
self.assertEqual(messages[0]["content"], "msg-5")
self.assertEqual(messages[-1]["content"], "msg-24")
DatabaseManager.reset_instance()
def test_file_sqlite_enables_wal_and_busy_timeout(self):
temp_dir = tempfile.TemporaryDirectory()
db_path = os.path.join(temp_dir.name, "sqlite_pragmas.db")
original_env = {
"DATABASE_PATH": os.environ.get("DATABASE_PATH"),
"SQLITE_BUSY_TIMEOUT_MS": os.environ.get("SQLITE_BUSY_TIMEOUT_MS"),
"SQLITE_WAL_ENABLED": os.environ.get("SQLITE_WAL_ENABLED"),
}
try:
os.environ["DATABASE_PATH"] = db_path
os.environ["SQLITE_BUSY_TIMEOUT_MS"] = "1234"
os.environ["SQLITE_WAL_ENABLED"] = "true"
Config.reset_instance()
DatabaseManager.reset_instance()
db = DatabaseManager.get_instance()
with db.get_session() as session:
journal_mode = session.connection().exec_driver_sql("PRAGMA journal_mode").scalar()
busy_timeout = session.connection().exec_driver_sql("PRAGMA busy_timeout").scalar()
self.assertEqual(str(journal_mode).lower(), "wal")
self.assertEqual(int(busy_timeout), 1234)
finally:
DatabaseManager.reset_instance()
Config.reset_instance()
for key, value in original_env.items():
if value is None:
os.environ.pop(key, None)
else:
os.environ[key] = value
temp_dir.cleanup()
def test_get_instance_waits_for_cold_start_initialization(self):
DatabaseManager.reset_instance()
Config.reset_instance()
temp_dir = tempfile.TemporaryDirectory()
db_path = os.path.join(temp_dir.name, "sqlite_cold_start.db")
original_database_path = os.environ.get("DATABASE_PATH")
create_all_entered = threading.Event()
competitor_entered = threading.Event()
release_create_all = threading.Event()
competitor_done = threading.Event()
state_lock = threading.Lock()
init_errors = []
competitor_errors = []
instances = []
query_values = []
original_create_all = Base.metadata.create_all
def delayed_create_all(bind, *args, **kwargs):
create_all_entered.set()
if not release_create_all.wait(timeout=5):
raise TimeoutError("Timed out waiting to release create_all")
return original_create_all(bind, *args, **kwargs)
def initialize_manager() -> None:
try:
db = DatabaseManager.get_instance()
with state_lock:
instances.append(db)
except Exception as exc:
with state_lock:
init_errors.append(exc)
def use_manager() -> None:
try:
competitor_entered.set()
db = DatabaseManager.get_instance()
session = db.get_session()
try:
value = session.connection().exec_driver_sql("SELECT 1").scalar()
finally:
session.close()
with state_lock:
instances.append(db)
query_values.append(value)
except Exception as exc:
with state_lock:
competitor_errors.append(exc)
finally:
competitor_done.set()
try:
os.environ["DATABASE_PATH"] = db_path
Config.reset_instance()
with patch.object(Base.metadata, "create_all", side_effect=delayed_create_all):
init_thread = threading.Thread(target=initialize_manager)
competitor_thread = threading.Thread(target=use_manager)
init_thread.start()
self.assertTrue(create_all_entered.wait(timeout=5))
competitor_thread.start()
self.assertTrue(competitor_entered.wait(timeout=5))
self.assertFalse(
competitor_done.wait(timeout=0.2),
"DatabaseManager.get_instance() returned before initialization completed",
)
release_create_all.set()
init_thread.join(timeout=5)
competitor_thread.join(timeout=5)
self.assertFalse(init_thread.is_alive())
self.assertFalse(competitor_thread.is_alive())
self.assertEqual(init_errors, [])
self.assertEqual(competitor_errors, [])
self.assertEqual(query_values, [1])
self.assertEqual(len({id(instance) for instance in instances}), 1)
finally:
release_create_all.set()
DatabaseManager.reset_instance()
Config.reset_instance()
if original_database_path is None:
os.environ.pop("DATABASE_PATH", None)
else:
os.environ["DATABASE_PATH"] = original_database_path
temp_dir.cleanup()
def test_direct_construction_serializes_before_get_instance(self):
DatabaseManager.reset_instance()
Config.reset_instance()
temp_dir = tempfile.TemporaryDirectory()
direct_db_path = os.path.join(temp_dir.name, "direct.db")
env_db_path = os.path.join(temp_dir.name, "env.db")
direct_db_url = f"sqlite:///{direct_db_path}"
original_database_path = os.environ.get("DATABASE_PATH")
direct_init_entered = threading.Event()
competitor_entered = threading.Event()
allow_direct_init = threading.Event()
competitor_done = threading.Event()
state_lock = threading.Lock()
errors = []
instances = []
query_values = []
original_init = DatabaseManager.__init__
def delayed_direct_init(self, db_url=None):
if db_url == direct_db_url:
direct_init_entered.set()
if not competitor_entered.wait(timeout=5):
raise TimeoutError("Timed out waiting for competitor")
if not allow_direct_init.wait(timeout=5):
raise TimeoutError("Timed out waiting to initialize direct instance")
return original_init(self, db_url=db_url)
def construct_directly() -> None:
try:
db = DatabaseManager(db_url=direct_db_url)
with state_lock:
instances.append(db)
except Exception as exc:
with state_lock:
errors.append(exc)
def use_get_instance() -> None:
try:
competitor_entered.set()
db = DatabaseManager.get_instance()
session = db.get_session()
try:
value = session.connection().exec_driver_sql("SELECT 1").scalar()
finally:
session.close()
with state_lock:
instances.append(db)
query_values.append(value)
except Exception as exc:
with state_lock:
errors.append(exc)
finally:
competitor_done.set()
try:
os.environ["DATABASE_PATH"] = env_db_path
Config.reset_instance()
with patch.object(DatabaseManager, "__init__", new=delayed_direct_init):
direct_thread = threading.Thread(target=construct_directly)
competitor_thread = threading.Thread(target=use_get_instance)
direct_thread.start()
self.assertTrue(direct_init_entered.wait(timeout=5))
competitor_thread.start()
self.assertTrue(competitor_entered.wait(timeout=5))
self.assertFalse(
competitor_done.wait(timeout=0.2),
"get_instance() should not initialize over an in-flight direct construction",
)
allow_direct_init.set()
direct_thread.join(timeout=5)
competitor_thread.join(timeout=5)
self.assertFalse(direct_thread.is_alive())
self.assertFalse(competitor_thread.is_alive())
self.assertEqual(errors, [])
self.assertEqual(query_values, [1])
self.assertEqual(len({id(instance) for instance in instances}), 1)
self.assertEqual(DatabaseManager._instance._db_url, direct_db_url)
finally:
allow_direct_init.set()
DatabaseManager.reset_instance()
Config.reset_instance()
if original_database_path is None:
os.environ.pop("DATABASE_PATH", None)
else:
os.environ["DATABASE_PATH"] = original_database_path
temp_dir.cleanup()
def test_init_cleanup_preserves_original_initialization_error(self):
DatabaseManager.reset_instance()
original_error = RuntimeError("create all failed")
cleanup_error = RuntimeError("dispose failed")
def create_engine_with_failing_dispose(*args, **kwargs):
engine = sqlalchemy_create_engine(*args, **kwargs)
def failing_dispose() -> None:
raise cleanup_error
engine.dispose = failing_dispose
return engine
try:
with patch("src.storage.create_engine", side_effect=create_engine_with_failing_dispose):
with patch.object(Base.metadata, "create_all", side_effect=original_error):
with self.assertRaisesRegex(RuntimeError, "create all failed") as ctx:
DatabaseManager.get_instance()
self.assertIs(ctx.exception, original_error)
self.assertIsNone(DatabaseManager._instance)
finally:
DatabaseManager.reset_instance()
def test_sqlite_write_transactions_begin_immediate(self):
DatabaseManager.reset_instance()
db = DatabaseManager(db_url="sqlite:///:memory:")
session = db.get_session()
connection = session.connection()
try:
with patch.object(db, "get_session", return_value=session):
with patch.object(connection, "exec_driver_sql", wraps=connection.exec_driver_sql) as mock_exec:
result = db._run_write_transaction("unit-test", lambda current_session: 7)
self.assertEqual(result, 7)
self.assertTrue(
any(call.args == ("BEGIN IMMEDIATE",) for call in mock_exec.call_args_list)
)
finally:
DatabaseManager.reset_instance()
def test_save_daily_data_sqlite_concurrent_same_code_date_counts_only_new_rows(self):
DatabaseManager.reset_instance()
temp_dir = tempfile.TemporaryDirectory()
db_path = os.path.join(temp_dir.name, "sqlite_daily_concurrency.db")
db = DatabaseManager(db_url=f"sqlite:///{db_path}")
results = []
results_lock = threading.Lock()
start_barrier = threading.Barrier(2)
def worker() -> None:
start_barrier.wait()
count = db.save_daily_data(
pd.DataFrame(
[
{
'date': date(2026, 4, 1),
'open': 10,
'high': 11,
'low': 9,
'close': 10.5,
'volume': 100,
'amount': 1050,
'pct_chg': 1.2,
'ma5': 10.1,
'ma10': 10.2,
'ma20': 10.3,
'volume_ratio': 1.0,
}
]
),
code='600519',
data_source='test',
)
with results_lock:
results.append(count)
threads = [threading.Thread(target=worker) for _ in range(2)]
for thread in threads:
thread.start()
for thread in threads:
thread.join()
try:
self.assertCountEqual(results, [1, 0])
with db.get_session() as session:
total = session.execute(
select(func.count()).select_from(StockDaily).where(
and_(
StockDaily.code == '600519',
StockDaily.date == date(2026, 4, 1),
)
)
).scalar()
self.assertEqual(total, 1)
finally:
temp_dir.cleanup()
DatabaseManager.reset_instance()
# ------------------------------------------------------------------
# Story 1.2 — canonical_id column + dual-write (issue #2207 PR2)
# ------------------------------------------------------------------
@staticmethod
def _create_legacy_stock_daily_without_canonical_id(db_path: str) -> None:
"""Pre-PR2 ``stock_daily`` schema: no ``canonical_id`` column."""
with sqlite3.connect(db_path) as conn:
conn.execute(
"""CREATE TABLE stock_daily (
id INTEGER PRIMARY KEY AUTOINCREMENT,
code VARCHAR(10) NOT NULL,
date DATE NOT NULL,
open FLOAT,
high FLOAT,
low FLOAT,
close FLOAT,
volume FLOAT,
amount FLOAT,
pct_chg FLOAT,
ma5 FLOAT,
ma10 FLOAT,
ma20 FLOAT,
volume_ratio FLOAT,
data_source VARCHAR(50),
created_at DATETIME,
updated_at DATETIME
)"""
)
conn.execute(
"CREATE UNIQUE INDEX uix_code_date ON stock_daily (code, date)"
)
conn.execute(
"CREATE INDEX ix_code_date ON stock_daily (code, date)"
)
@staticmethod
def _create_legacy_stock_daily_with_canonical_id(db_path: str) -> None:
"""``stock_daily`` schema WITH a ``canonical_id`` column (for repair tests)."""
with sqlite3.connect(db_path) as conn:
conn.execute(
"""CREATE TABLE stock_daily (
id INTEGER PRIMARY KEY AUTOINCREMENT,
code VARCHAR(10) NOT NULL,
date DATE NOT NULL,
open FLOAT,
high FLOAT,
low FLOAT,
close FLOAT,
volume FLOAT,
amount FLOAT,
pct_chg FLOAT,
ma5 FLOAT,
ma10 FLOAT,
ma20 FLOAT,
volume_ratio FLOAT,
data_source VARCHAR(50),
created_at DATETIME,
updated_at DATETIME,
canonical_id VARCHAR(32)
)"""
)
conn.execute(
"CREATE UNIQUE INDEX uix_code_date ON stock_daily (code, date)"
)
conn.execute(
"CREATE INDEX ix_code_date ON stock_daily (code, date)"
)
@staticmethod
def _make_temp_db_path() -> tuple:
"""Return (db_dir, db_path); Windows-safe cleanup via shutil.rmtree."""
db_dir = tempfile.mkdtemp(prefix="dsa_canonical_id_")
db_path = os.path.join(db_dir, "test.db")
return db_dir, db_path
@staticmethod
def _cleanup_temp_dir(db_dir: str) -> None:
# Windows: the SQLAlchemy engine pool may still hold the SQLite file
# handle briefly after dispose(); ignore_errors lets teardown succeed.
shutil.rmtree(db_dir, ignore_errors=True)
def test_canonical_id_migration_adds_column_and_backfills_existing_rows(self):
"""AC 1: ALTER adds ``canonical_id``; existing rows backfilled via parser."""
DatabaseManager.reset_instance()
db_dir, db_path = self._make_temp_db_path()
try:
self._create_legacy_stock_daily_without_canonical_id(db_path)
with sqlite3.connect(db_path) as conn:
conn.executemany(
"""INSERT INTO stock_daily (
code, date, close, data_source
) VALUES (?, ?, ?, ?)""",
[
("sh000300", "2026-01-01", 4000.0, "legacy"),
("600519", "2026-01-01", 1600.0, "legacy"),
("AAPL", "2026-01-01", 150.0, "legacy"),
],
)
DatabaseManager(db_url=f"sqlite:///{db_path}")
with sqlite3.connect(db_path) as conn:
columns = {row[1] for row in conn.execute("PRAGMA table_info(stock_daily)").fetchall()}
rows = {
code: canonical_id
for code, canonical_id in conn.execute(
"SELECT code, canonical_id FROM stock_daily ORDER BY id"
).fetchall()
}
self.assertIn("canonical_id", columns)
# sh000300 is a registered index alias → sh000300
self.assertEqual(rows["sh000300"], "sh000300")
# 600519 bare A-share → sh600519 (contract: bare codes default stock,
# classifier routes 6-prefixed to SH)
self.assertEqual(rows["600519"], "sh600519")
# AAPL bare US ticker → AAPL (canonical US form, no prefix)
self.assertEqual(rows["AAPL"], "AAPL")
finally:
DatabaseManager.reset_instance()
Config.reset_instance()
self._cleanup_temp_dir(db_dir)
def test_canonical_id_migration_is_idempotent(self):
"""AC 2: second startup detects the column and skips ALTER; backfill is a no-op."""
DatabaseManager.reset_instance()
db_dir, db_path = self._make_temp_db_path()
try:
self._create_legacy_stock_daily_without_canonical_id(db_path)
with sqlite3.connect(db_path) as conn:
conn.execute(
"INSERT INTO stock_daily (code, date, close) VALUES (?, ?, ?)",
("600519", "2026-01-01", 1600.0),
)
# First run: adds column + backfills.
DatabaseManager(db_url=f"sqlite:///{db_path}")
with sqlite3.connect(db_path) as conn:
first_pass = conn.execute(
"SELECT canonical_id FROM stock_daily WHERE code='600519'"
).fetchone()[0]
self.assertEqual(first_pass, "sh600519")
DatabaseManager.reset_instance()
# Second run: must not raise, must not duplicate the column.
DatabaseManager(db_url=f"sqlite:///{db_path}")
with sqlite3.connect(db_path) as conn:
columns = [row[1] for row in conn.execute("PRAGMA table_info(stock_daily)").fetchall()]
second_pass = conn.execute(
"SELECT canonical_id FROM stock_daily WHERE code='600519'"
).fetchone()[0]
self.assertEqual(columns.count("canonical_id"), 1)
self.assertEqual(second_pass, "sh600519")
finally:
DatabaseManager.reset_instance()
Config.reset_instance()
self._cleanup_temp_dir(db_dir)
def test_canonical_id_migration_on_fresh_empty_db(self):
"""AC: empty/new DB has the column from create_all; backfill finds 0 rows."""
DatabaseManager.reset_instance()
db_dir, db_path = self._make_temp_db_path()
try:
DatabaseManager(db_url=f"sqlite:///{db_path}")
with sqlite3.connect(db_path) as conn:
columns = {row[1] for row in conn.execute("PRAGMA table_info(stock_daily)").fetchall()}
count = conn.execute("SELECT COUNT(*) FROM stock_daily").fetchone()[0]
self.assertIn("canonical_id", columns)
self.assertEqual(count, 0)
indexes = self._list_sqlite_indexes(db_path, "stock_daily")
self.assertEqual(indexes.get("ix_stock_daily_canonical_id"), ["canonical_id"])
finally:
DatabaseManager.reset_instance()
Config.reset_instance()
self._cleanup_temp_dir(db_dir)
def test_canonical_id_migration_creates_plain_non_unique_index(self):
"""AC 6/9: index is plain (not unique) so alias rows can coexist."""
DatabaseManager.reset_instance()
db_dir, db_path = self._make_temp_db_path()
try:
DatabaseManager(db_url=f"sqlite:///{db_path}")
unique_indexes = self._list_sqlite_unique_indexes(db_path, "stock_daily")
self.assertNotIn("ix_stock_daily_canonical_id", unique_indexes)
indexes = self._list_sqlite_indexes(db_path, "stock_daily")
self.assertEqual(indexes.get("ix_stock_daily_canonical_id"), ["canonical_id"])
finally:
DatabaseManager.reset_instance()
Config.reset_instance()
self._cleanup_temp_dir(db_dir)
def test_canonical_id_backfill_bare_code_colliding_with_index_stays_stock(self):
"""``_derive_canonical_id`` no longer reads
``matched_index.canonical_id``. A bare ``000300`` (which collides with
the CSI-300 index) now derives to the stock-path canonical_id
``sz000300`` — the parser contract says bare codes are always stock.
The index conflict is surfaced via ``matched_index`` only, never used
to override the canonical_id."""
DatabaseManager.reset_instance()
db_dir, db_path = self._make_temp_db_path()
try:
self._create_legacy_stock_daily_without_canonical_id(db_path)
with sqlite3.connect(db_path) as conn:
conn.execute(
"INSERT INTO stock_daily (code, date, close) VALUES (?, ?, ?)",
("000300", "2026-01-01", 4000.0),
)
DatabaseManager(db_url=f"sqlite:///{db_path}")
with sqlite3.connect(db_path) as conn:
canonical_id = conn.execute(
"SELECT canonical_id FROM stock_daily WHERE code='000300'"
).fetchone()[0]
# Bare code resolves to the stock-path canonical_id.
self.assertEqual(canonical_id, "sz000300")
finally:
DatabaseManager.reset_instance()
Config.reset_instance()
self._cleanup_temp_dir(db_dir)
def test_canonical_id_backfill_bare_non_index_code_stays_stock(self):
"""Index-aware backfill does not affect bare non-index codes: ``600519``
has no registry hit (``matched_index is None``) and backfills to the
stock-path canonical_id ``sh600519`` (contract #2 — bare → stock)."""
DatabaseManager.reset_instance()
db_dir, db_path = self._make_temp_db_path()
try:
self._create_legacy_stock_daily_without_canonical_id(db_path)
with sqlite3.connect(db_path) as conn:
conn.execute(
"INSERT INTO stock_daily (code, date, close) VALUES (?, ?, ?)",
("600519", "2026-01-01", 1600.0),
)
DatabaseManager(db_url=f"sqlite:///{db_path}")
with sqlite3.connect(db_path) as conn:
canonical_id = conn.execute(
"SELECT canonical_id FROM stock_daily WHERE code='600519'"
).fetchone()[0]
self.assertEqual(canonical_id, "sh600519")
finally:
DatabaseManager.reset_instance()
Config.reset_instance()
self._cleanup_temp_dir(db_dir)
def test_canonical_id_backfill_preserves_duplicate_alias_rows(self):
"""AC 9: multiple rows sharing canonical_id + date coexist (plain index)."""
DatabaseManager.reset_instance()
db_dir, db_path = self._make_temp_db_path()
try:
self._create_legacy_stock_daily_without_canonical_id(db_path)
with sqlite3.connect(db_path) as conn:
# Two distinct code aliases for the same underlying, same date.
conn.executemany(
"""INSERT INTO stock_daily (code, date, close) VALUES (?, ?, ?)""",
[
("sh000300", "2026-01-01", 4000.0),
("000300.SH", "2026-01-01", 4001.0),
],
)
DatabaseManager(db_url=f"sqlite:///{db_path}")
with sqlite3.connect(db_path) as conn:
rows = conn.execute(
"SELECT code, canonical_id FROM stock_daily ORDER BY id"
).fetchall()
self.assertEqual(len(rows), 2)
# Both survive — the plain index allows duplicate canonical_ids.
canonical_ids = {row[1] for row in rows}
self.assertEqual(canonical_ids, {"sh000300"})
finally:
DatabaseManager.reset_instance()
Config.reset_instance()
self._cleanup_temp_dir(db_dir)
def test_canonical_id_backfill_survives_per_row_derivation_failure(self):
"""D1 degrade: one unparseable row must not trap the batch scan loop.
The backfill advances an id cursor, so a row whose derivation raises
is skipped past (left NULL, retried next startup) instead of being
re-selected forever. Good rows after the bad one must still backfill.
"""
real_parse = stock_list_parser_module.parse_analysis_target
def flaky_parse(code):
if code == "!!broken!!":
raise RuntimeError("parser boom")
return real_parse(code)
DatabaseManager.reset_instance()
db_dir, db_path = self._make_temp_db_path()
try:
self._create_legacy_stock_daily_without_canonical_id(db_path)
with sqlite3.connect(db_path) as conn:
# Bad row FIRST, good rows after — maximizes re-select exposure
# if the scan ever regresses to "re-select NULL rows".
conn.executemany(
"""INSERT INTO stock_daily (code, date, close) VALUES (?, ?, ?)""",
[
("!!broken!!", "2026-01-01", 1.0),
("sh600519", "2026-01-01", 1700.0),
("AAPL", "2026-01-01", 210.0),
],
)
with patch(
"src.services.stock_list_parser.parse_analysis_target",
side_effect=flaky_parse,
):
DatabaseManager(db_url=f"sqlite:///{db_path}")
with sqlite3.connect(db_path) as conn:
rows = dict(
conn.execute(
"SELECT code, canonical_id FROM stock_daily"
).fetchall()
)
self.assertIsNone(rows["!!broken!!"])
self.assertEqual(rows["sh600519"], "sh600519")
self.assertEqual(rows["AAPL"], "AAPL")
finally:
DatabaseManager.reset_instance()
Config.reset_instance()
self._cleanup_temp_dir(db_dir)
def test_canonical_id_backfill_skips_empty_derived_values(self):
"""Review fix: parser returning None/'' (without raising) must not
count as backfilled nor persist an empty string as a stable key."""
real_parse = stock_list_parser_module.parse_analysis_target
def empty_parse(code):
if code == "!!empty!!":
return None
return real_parse(code)
DatabaseManager.reset_instance()
db_dir, db_path = self._make_temp_db_path()
try:
self._create_legacy_stock_daily_without_canonical_id(db_path)
with sqlite3.connect(db_path) as conn:
conn.executemany(
"""INSERT INTO stock_daily (code, date, close) VALUES (?, ?, ?)""",
[
("!!empty!!", "2026-01-01", 1.0),
("sh600519", "2026-01-01", 1700.0),
],
)
with patch(
"src.services.stock_list_parser.parse_analysis_target",
side_effect=empty_parse,
):
DatabaseManager(db_url=f"sqlite:///{db_path}")
with sqlite3.connect(db_path) as conn:
rows = dict(
conn.execute(
"SELECT code, canonical_id FROM stock_daily"
).fetchall()
)
self.assertIsNone(rows["!!empty!!"])
self.assertEqual(rows["sh600519"], "sh600519")
finally:
DatabaseManager.reset_instance()
Config.reset_instance()
self._cleanup_temp_dir(db_dir)
def test_save_daily_data_dual_writes_explicit_canonical_id(self):
"""AC 3: explicit ``canonical_id`` is written to both columns on upsert."""
DatabaseManager.reset_instance()
db_dir, db_path = self._make_temp_db_path()
db = DatabaseManager(db_url=f"sqlite:///{db_path}")
try:
df = pd.DataFrame(
[
{
'date': date(2026, 4, 1),
'open': 10, 'high': 11, 'low': 9, 'close': 10.5,
'volume': 100, 'amount': 1050, 'pct_chg': 1.2,
'ma5': 10.1, 'ma10': 10.2, 'ma20': 10.3, 'volume_ratio': 1.0,
}
]
)
db.save_daily_data(
df,
code="sh000016",
data_source="tencent",
canonical_id="sh000016",
)
with db.get_session() as session:
row = session.execute(
select(StockDaily).where(
and_(StockDaily.code == "sh000016", StockDaily.date == date(2026, 4, 1))
)
).scalar_one()
self.assertEqual(row.code, "sh000016")
self.assertEqual(row.canonical_id, "sh000016")
# Upsert path: re-saving updates canonical_id too.
db.save_daily_data(
df,
code="sh000016",
data_source="tencent",
canonical_id="sh000016-v2",
)
with db.get_session() as session:
row = session.execute(
select(StockDaily).where(
and_(StockDaily.code == "sh000016", StockDaily.date == date(2026, 4, 1))
)
).scalar_one()
self.assertEqual(row.canonical_id, "sh000016-v2")
finally:
DatabaseManager.reset_instance()
self._cleanup_temp_dir(db_dir)
def test_save_daily_data_derives_canonical_id_when_omitted(self):
"""D1/AC: ``canonical_id`` omitted → auto-derive via parser; parser failure → NULL, no raise."""
DatabaseManager.reset_instance()
db_dir, db_path = self._make_temp_db_path()
db = DatabaseManager(db_url=f"sqlite:///{db_path}")
try:
df = pd.DataFrame(
[
{
'date': date(2026, 4, 1),
'open': 10, 'high': 11, 'low': 9, 'close': 10.5,
'volume': 100, 'amount': 1050, 'pct_chg': 1.2,
'ma5': 10.1, 'ma10': 10.2, 'ma20': 10.3, 'volume_ratio': 1.0,
}
]
)
# Main path: bare A-share → sh600519.
db.save_daily_data(df, code="600519", data_source="test")
with db.get_session() as session:
row = session.execute(
select(StockDaily).where(
and_(StockDaily.code == "600519", StockDaily.date == date(2026, 4, 1))
)
).scalar_one()
self.assertEqual(row.canonical_id, "sh600519")
# Degraded path: parser failure on re-save must NOT overwrite a
# previously backfilled non-NULL canonical_id with NULL (review
# fix: coalesce preserves the existing value). No exception
# bubbles out of save_daily_data (D1).
with patch(
"src.services.stock_list_parser.parse_analysis_target",
side_effect=RuntimeError("parser boom"),
):
db.save_daily_data(df, code="600519", data_source="test")
with db.get_session() as session:
row = session.execute(
select(StockDaily).where(
and_(StockDaily.code == "600519", StockDaily.date == date(2026, 4, 1))
)
).scalar_one()
self.assertEqual(row.canonical_id, "sh600519")
finally:
DatabaseManager.reset_instance()
self._cleanup_temp_dir(db_dir)
def test_save_daily_data_empty_string_canonical_id_triggers_derivation(self):
"""Review fix: empty-string canonical_id is treated like None —
derived via parser, never persisted as an empty key."""
DatabaseManager.reset_instance()
db_dir, db_path = self._make_temp_db_path()
db = DatabaseManager(db_url=f"sqlite:///{db_path}")
try:
df = pd.DataFrame(
[
{
'date': date(2026, 4, 3),
'open': 10, 'high': 11, 'low': 9, 'close': 10.5,
'volume': 100, 'amount': 1050, 'pct_chg': 1.2,
'ma5': 10.1, 'ma10': 10.2, 'ma20': 10.3, 'volume_ratio': 1.0,
}
]
)
db.save_daily_data(df, code="600519", data_source="test", canonical_id="")
with db.get_session() as session:
row = session.execute(
select(StockDaily).where(
and_(StockDaily.code == "600519", StockDaily.date == date(2026, 4, 3))
)
).scalar_one()
self.assertEqual(row.canonical_id, "sh600519")
finally:
DatabaseManager.reset_instance()
self._cleanup_temp_dir(db_dir)
def test_save_daily_data_first_insert_parser_failure_writes_null(self):
"""D1: first insert with parser failure → row created with NULL canonical_id, no raise."""
DatabaseManager.reset_instance()
db_dir, db_path = self._make_temp_db_path()
db = DatabaseManager(db_url=f"sqlite:///{db_path}")
try:
df = pd.DataFrame(
[
{
'date': date(2026, 4, 2),
'open': 10, 'high': 11, 'low': 9, 'close': 10.5,
'volume': 100, 'amount': 1050, 'pct_chg': 1.2,
'ma5': 10.1, 'ma10': 10.2, 'ma20': 10.3, 'volume_ratio': 1.0,
}
]
)
with patch(
"src.services.stock_list_parser.parse_analysis_target",
side_effect=RuntimeError("parser boom"),
):
db.save_daily_data(df, code="600519", data_source="test")
with db.get_session() as session:
row = session.execute(
select(StockDaily).where(
and_(StockDaily.code == "600519", StockDaily.date == date(2026, 4, 2))
)
).scalar_one()
self.assertIsNone(row.canonical_id)
finally:
DatabaseManager.reset_instance()
self._cleanup_temp_dir(db_dir)
def test_save_daily_data_derives_stock_canonical_id_for_bare_index_collision(self):
"""``save_daily_data(df, code="000300")`` with no explicit
canonical_id writes ``sz000300`` (stock-path canonical_id), NOT
``sh000300`` (the index canonical_id). Bare codes are always stock."""
DatabaseManager.reset_instance()
db_dir, db_path = self._make_temp_db_path()
db = DatabaseManager(db_url=f"sqlite:///{db_path}")
try:
df = pd.DataFrame(
[
{
'date': date(2026, 4, 5),
'open': 10, 'high': 11, 'low': 9, 'close': 10.5,
'volume': 100, 'amount': 1050, 'pct_chg': 1.2,
'ma5': 10.1, 'ma10': 10.2, 'ma20': 10.3, 'volume_ratio': 1.0,
}
]
)
db.save_daily_data(df, code="000300", data_source="test")
with db.get_session() as session:
row = session.execute(
select(StockDaily).where(
and_(StockDaily.code == "000300", StockDaily.date == date(2026, 4, 5))
)
).scalar_one()
self.assertEqual(row.canonical_id, "sz000300")
finally:
DatabaseManager.reset_instance()
self._cleanup_temp_dir(db_dir)
def test_save_daily_data_explicit_index_forms_derive_index_canonical_id(self):
"""Explicit index forms (``sh000300`` / ``000300.SH``)
derive to the index canonical_id ``sh000300``, while the bare code
``000300`` derives to the stock-path ``sz000300``. Explicit index and
bare stock are intentionally different buckets."""
DatabaseManager.reset_instance()
db_dir, db_path = self._make_temp_db_path()
db = DatabaseManager(db_url=f"sqlite:///{db_path}")
try:
df = pd.DataFrame(
[
{
'date': date(2026, 4, 6),
'open': 10, 'high': 11, 'low': 9, 'close': 10.5,
'volume': 100, 'amount': 1050, 'pct_chg': 1.2,
'ma5': 10.1, 'ma10': 10.2, 'ma20': 10.3, 'volume_ratio': 1.0,
}
]
)
for code in ("000300", "sh000300", "000300.SH"):
db.save_daily_data(df, code=code, data_source="test")
with db.get_session() as session:
rows = session.execute(
select(StockDaily).order_by(StockDaily.id)
).scalars().all()
self.assertEqual(len(rows), 3)
self.assertEqual(
{row.code: row.canonical_id for row in rows},
{
"000300": "sz000300",
"sh000300": "sh000300",
"000300.SH": "sh000300",
},
)
finally:
DatabaseManager.reset_instance()
self._cleanup_temp_dir(db_dir)
def test_canonical_id_migration_raises_when_column_inspection_fails(self):
"""AC: inspect() failure on stock_daily must raise (no silent downgrade).
Mirrors the decision_profile inspection-failure test (L417-438). The
patched inspector only breaks for ``stock_daily`` so the earlier
``decision_signal_profile`` migration (which runs first in the
``__init__`` ensure chain) still passes and execution reaches the
``_ensure_stock_daily_canonical_id`` step.
"""
real_inspector = inspect
class BrokenInspector:
def __init__(self, engine):
self._real = real_inspector(engine)
def has_table(self, table_name: str) -> bool:
return self._real.has_table(table_name)
def get_columns(self, table_name: str):
if table_name == StockDaily.__tablename__:
raise RuntimeError("inspection failed")
return self._real.get_columns(table_name)
def get_indexes(self, table_name: str):
return self._real.get_indexes(table_name)
DatabaseManager.reset_instance()
try:
with patch("src.storage.inspect", side_effect=BrokenInspector):
with self.assertLogs("src.storage", level="ERROR") as logs:
with self.assertRaises(RuntimeError):
DatabaseManager(db_url="sqlite:///:memory:")
self.assertIn(
"canonical_id migration cannot continue safely",
"\n".join(logs.output),
)
finally:
DatabaseManager.reset_instance()
Config.reset_instance()
# ------------------------------------------------------------------
# canonical_id repair
# ------------------------------------------------------------------
def test_derive_canonical_id_bare_conflict_stays_stock(self):
"""``_derive_canonical_id`` returns the parser stock
canonical for bare conflict codes, never the index canonical."""
DatabaseManager.reset_instance()
db = DatabaseManager(db_url="sqlite:///:memory:")
try:
self.assertEqual(db._derive_canonical_id("000016"), "sz000016")
self.assertEqual(db._derive_canonical_id("930955"), "bj930955")
self.assertEqual(db._derive_canonical_id("sh000016"), "sh000016")
self.assertEqual(db._derive_canonical_id("930955.CSI"), "csi930955")
finally:
DatabaseManager.reset_instance()
def test_derive_canonical_id_rejects_unregistered_csi(self):
"""PR #2267 review fix: an unsupported identity (unregistered
``csi`` prefix or ``.CSI`` suffix) must NOT enter a persistent
canonical bucket — ``_derive_canonical_id`` returns None so the
caller persists NULL instead of the raw token."""
DatabaseManager.reset_instance()
db = DatabaseManager(db_url="sqlite:///:memory:")
try:
self.assertIsNone(db._derive_canonical_id("csi930956"))
self.assertIsNone(db._derive_canonical_id("CSI930956"))
self.assertIsNone(db._derive_canonical_id("930956.CSI"))
self.assertIsNone(db._derive_canonical_id("csi000300"))
finally:
DatabaseManager.reset_instance()
def test_canonical_id_repair_fixes_bare_misbucketed_rows(self):
"""Rows whose bare code has an erroneous index canonical are
repaired to the parser stock canonical; explicit index rows and
correct stock rows are untouched."""
DatabaseManager.reset_instance()
db_dir, db_path = self._make_temp_db_path()
try:
self._create_legacy_stock_daily_with_canonical_id(db_path)
with sqlite3.connect(db_path) as conn:
conn.executemany(
"""INSERT INTO stock_daily (code, date, close, canonical_id)
VALUES (?, ?, ?, ?)""",
[
# Mis-bucketed bare codes (old index-aware derivation).
("000001", "2026-01-01", 1.0, "sh000001"),
("000016", "2026-01-02", 2.0, "sh000016"),
("000688", "2026-01-03", 3.0, "sh000688"),
("930955", "2026-01-04", 4.0, "csi930955"),
# Explicit index rows — must NOT be modified.
("sh000016", "2026-01-05", 5.0, "sh000016"),
("930955.CSI", "2026-01-06", 6.0, "csi930955"),
# Correct stock row — must NOT be modified.
("600519", "2026-01-07", 7.0, "sh600519"),
# Unrelated row — must NOT be modified.
("AAPL", "2026-01-08", 8.0, "AAPL"),
],
)
DatabaseManager(db_url=f"sqlite:///{db_path}")
with sqlite3.connect(db_path) as conn:
rows = dict(
conn.execute(
"SELECT code, canonical_id FROM stock_daily ORDER BY id"
).fetchall()
)
self.assertEqual(rows["000001"], "sz000001")
self.assertEqual(rows["000016"], "sz000016")
self.assertEqual(rows["000688"], "sz000688")
self.assertEqual(rows["930955"], "bj930955")
# Explicit index rows preserved.
self.assertEqual(rows["sh000016"], "sh000016")
self.assertEqual(rows["930955.CSI"], "csi930955")
# Correct stock + unrelated rows preserved.
self.assertEqual(rows["600519"], "sh600519")
self.assertEqual(rows["AAPL"], "AAPL")
finally:
DatabaseManager.reset_instance()
Config.reset_instance()
self._cleanup_temp_dir(db_dir)
def test_canonical_id_repair_is_idempotent(self):
"""Running the repair a second time repairs 0 rows."""
DatabaseManager.reset_instance()
db_dir, db_path = self._make_temp_db_path()
try:
self._create_legacy_stock_daily_with_canonical_id(db_path)
with sqlite3.connect(db_path) as conn:
conn.execute(
"INSERT INTO stock_daily (code, date, close, canonical_id) "
"VALUES (?, ?, ?, ?)",
("000016", "2026-01-01", 2.0, "sh000016"),
)
DatabaseManager(db_url=f"sqlite:///{db_path}")
with sqlite3.connect(db_path) as conn:
first = conn.execute(
"SELECT canonical_id FROM stock_daily WHERE code='000016'"
).fetchone()[0]
self.assertEqual(first, "sz000016")
DatabaseManager.reset_instance()
with self.assertLogs("src.storage", level="INFO") as logs:
DatabaseManager(db_url=f"sqlite:///{db_path}")
log_text = "\n".join(logs.output)
self.assertIn("repaired_count=0", log_text)
finally:
DatabaseManager.reset_instance()
Config.reset_instance()
self._cleanup_temp_dir(db_dir)
def test_canonical_id_repair_skips_when_registry_empty(self):
"""When the index registry is empty the repair is a no-op and
logs a WARNING; no rows are modified."""
DatabaseManager.reset_instance()
db_dir, db_path = self._make_temp_db_path()
try:
self._create_legacy_stock_daily_with_canonical_id(db_path)
with sqlite3.connect(db_path) as conn:
conn.execute(
"INSERT INTO stock_daily (code, date, close, canonical_id) "
"VALUES (?, ?, ?, ?)",
("000016", "2026-01-01", 2.0, "sh000016"),
)
with patch(
"src.data.stock_index_loader._load_active_index_rows",
return_value=[],
), self.assertLogs("src.storage", level="WARNING") as logs:
DatabaseManager(db_url=f"sqlite:///{db_path}")
with sqlite3.connect(db_path) as conn:
canonical_id = conn.execute(
"SELECT canonical_id FROM stock_daily WHERE code='000016'"
).fetchone()[0]
self.assertEqual(canonical_id, "sh000016")
self.assertTrue(
any("registry is empty" in record.getMessage() for record in logs.records)
)
finally:
DatabaseManager.reset_instance()
Config.reset_instance()
self._cleanup_temp_dir(db_dir)
def test_canonical_id_repair_handles_concurrent_rewrite_safely(self):
"""A conditional UPDATE that loses the race (row already
rewritten) is skipped, not double-counted."""
DatabaseManager.reset_instance()
db_dir, db_path = self._make_temp_db_path()
try:
self._create_legacy_stock_daily_with_canonical_id(db_path)
with sqlite3.connect(db_path) as conn:
conn.execute(
"INSERT INTO stock_daily (code, date, close, canonical_id) "
"VALUES (?, ?, ?, ?)",
("000016", "2026-01-01", 2.0, "sh000016"),
)
real_parse = stock_list_parser_module.parse_analysis_target
def racing_parse(code):
# Simulate a concurrent writer that already fixed the row.
if code == "000016":
with sqlite3.connect(db_path) as conn:
conn.execute(
"UPDATE stock_daily SET canonical_id='sz000016' "
"WHERE code='000016'"
)
conn.commit()
return real_parse(code)
with patch(
"src.services.stock_list_parser.parse_analysis_target",
side_effect=racing_parse,
):
DatabaseManager(db_url=f"sqlite:///{db_path}")
with sqlite3.connect(db_path) as conn:
canonical_id = conn.execute(
"SELECT canonical_id FROM stock_daily WHERE code='000016'"
).fetchone()[0]
self.assertEqual(canonical_id, "sz000016")
finally:
DatabaseManager.reset_instance()
Config.reset_instance()
self._cleanup_temp_dir(db_dir)
def test_canonical_id_repair_queries_only_relevant_rows(self):
"""Gap 6: the repair SQL filters to rows whose canonical_id is in the
active index set, so rows with unrelated canonical_ids are never
scanned/derived."""
DatabaseManager.reset_instance()
db_dir, db_path = self._make_temp_db_path()
try:
self._create_legacy_stock_daily_with_canonical_id(db_path)
with sqlite3.connect(db_path) as conn:
conn.executemany(
"""INSERT INTO stock_daily (code, date, close, canonical_id)
VALUES (?, ?, ?, ?)""",
[
# Mis-bucketed bare code — must be repaired.
("000016", "2026-01-01", 2.0, "sh000016"),
# Unrelated canonical — must NOT be touched.
("600519", "2026-01-02", 7.0, "sh600519"),
("AAPL", "2026-01-03", 8.0, "AAPL"),
# Explicit index row — must NOT be touched.
("sh000016", "2026-01-04", 5.0, "sh000016"),
],
)
DatabaseManager(db_url=f"sqlite:///{db_path}")
with sqlite3.connect(db_path) as conn:
rows = dict(
conn.execute(
"SELECT code, canonical_id FROM stock_daily ORDER BY id"
).fetchall()
)
# Only the mis-bucketed bare code is repaired.
self.assertEqual(rows["000016"], "sz000016")
self.assertEqual(rows["600519"], "sh600519")
self.assertEqual(rows["AAPL"], "AAPL")
self.assertEqual(rows["sh000016"], "sh000016")
finally:
DatabaseManager.reset_instance()
Config.reset_instance()
self._cleanup_temp_dir(db_dir)
if __name__ == '__main__':
unittest.main()