54 lines
2.2 KiB
Python
54 lines
2.2 KiB
Python
from __future__ import annotations
|
|
|
|
import tempfile
|
|
import unittest
|
|
from pathlib import Path
|
|
|
|
from backend.bootstrap import build_application_container
|
|
from backend.bootstrap.settings import environment_credentials
|
|
from database import ReviewDatabase
|
|
|
|
|
|
class BootstrapContainerTests(unittest.TestCase):
|
|
def test_environment_credentials_preserve_legacy_model_fallbacks(self) -> None:
|
|
result = environment_credentials(
|
|
{
|
|
"TUSHARE_TOKEN": " tushare ",
|
|
"IFIND_REFRESH_TOKEN": " refresh ",
|
|
"LLM_API_KEY": "legacy-key",
|
|
"LLM_BASE_URL": "https://legacy.example/v1",
|
|
"LLM_MODEL": "legacy-model",
|
|
}
|
|
)
|
|
self.assertEqual(result["tushare_token"], "tushare")
|
|
self.assertEqual(result["ifind_refresh_token"], "refresh")
|
|
self.assertEqual(result["platform_llm_primary_api_key"], "legacy-key")
|
|
self.assertEqual(result["platform_llm_primary_base_url"], "https://legacy.example/v1")
|
|
self.assertEqual(result["platform_llm_primary_model"], "legacy-model")
|
|
|
|
def test_container_shares_one_database_and_one_ifind_client(self) -> None:
|
|
with tempfile.TemporaryDirectory() as temporary:
|
|
root = Path(temporary)
|
|
public_skills = root / "public"
|
|
private_skills = root / "private"
|
|
public_skills.mkdir()
|
|
private_skills.mkdir()
|
|
database = ReviewDatabase(root / "review.db")
|
|
container = build_application_container(
|
|
database,
|
|
{"ifind_refresh_token": "refresh-token", "ifind_access_token": "access-token"},
|
|
public_skills,
|
|
private_skills,
|
|
)
|
|
self.assertIs(container.database, database)
|
|
self.assertIs(container.screener.database, database)
|
|
self.assertIs(container.strategy_tracking.database, database)
|
|
self.assertIs(container.alert_service.database, database)
|
|
self.assertIs(container.trade_journal.database, database)
|
|
self.assertIs(container.chart_data.ifind, container.ifind)
|
|
self.assertTrue(container.ifind.configured)
|
|
|
|
|
|
if __name__ == "__main__":
|
|
unittest.main()
|