diff --git a/backend/app/services/simulation_manager.py b/backend/app/services/simulation_manager.py index bda876cd..ff63d020 100644 --- a/backend/app/services/simulation_manager.py +++ b/backend/app/services/simulation_manager.py @@ -496,12 +496,24 @@ class SimulationManager: if platform is None: platform = state.get_default_platform() + if platform not in {"twitter", "reddit"}: + raise ValueError(f"不支持的平台: {platform}") + sim_dir = self._get_simulation_dir(simulation_id) - profile_path = os.path.join(sim_dir, f"{platform}_profiles.json") + profile_path = os.path.join( + sim_dir, + "twitter_profiles.csv" if platform == "twitter" else "reddit_profiles.json", + ) if not os.path.exists(profile_path): return [] - + + if platform == "twitter": + import csv + + with open(profile_path, 'r', encoding='utf-8', newline='') as f: + return list(csv.DictReader(f)) + with open(profile_path, 'r', encoding='utf-8') as f: return json.load(f) diff --git a/backend/tests/test_platform_profiles.py b/backend/tests/test_platform_profiles.py new file mode 100644 index 00000000..f3883bda --- /dev/null +++ b/backend/tests/test_platform_profiles.py @@ -0,0 +1,55 @@ +import csv +import json + +import pytest + +from app.services.simulation_manager import SimulationManager, SimulationState + + +def _manager(tmp_path, *, twitter, reddit): + manager = SimulationManager() + manager.SIMULATION_DATA_DIR = str(tmp_path) + state = SimulationState( + simulation_id="sim_test", + project_id="proj_test", + graph_id="graph_test", + enable_twitter=twitter, + enable_reddit=reddit, + ) + manager._save_simulation_state(state) + return manager, tmp_path / "sim_test" + + +def test_twitter_default_reads_generated_csv(tmp_path): + manager, sim_dir = _manager(tmp_path, twitter=True, reddit=False) + with (sim_dir / "twitter_profiles.csv").open("w", encoding="utf-8", newline="") as handle: + writer = csv.DictWriter(handle, fieldnames=["agent_id", "name"]) + writer.writeheader() + writer.writerow({"agent_id": "0", "name": "Alice"}) + + assert manager.get_profiles("sim_test") == [{"agent_id": "0", "name": "Alice"}] + + +def test_reddit_default_preserves_json_behavior(tmp_path): + manager, sim_dir = _manager(tmp_path, twitter=False, reddit=True) + profiles = [{"user_id": 0, "username": "alice"}] + (sim_dir / "reddit_profiles.json").write_text(json.dumps(profiles), encoding="utf-8") + + assert manager.get_profiles("sim_test") == profiles + + +def test_explicit_platform_overrides_dual_platform_default(tmp_path): + manager, sim_dir = _manager(tmp_path, twitter=True, reddit=True) + with (sim_dir / "twitter_profiles.csv").open("w", encoding="utf-8", newline="") as handle: + writer = csv.DictWriter(handle, fieldnames=["agent_id", "name"]) + writer.writeheader() + writer.writerow({"agent_id": "1", "name": "Bob"}) + + assert manager.get_profiles("sim_test", platform="twitter")[0]["name"] == "Bob" + + +def test_invalid_platform_is_rejected(tmp_path): + manager, _ = _manager(tmp_path, twitter=True, reddit=True) + + with pytest.raises(ValueError, match="不支持的平台"): + manager.get_profiles("sim_test", platform="../state")