fix: load Twitter profiles from CSV
This commit is contained in:
parent
aaf3ea0a09
commit
b651883d66
|
|
@ -496,12 +496,24 @@ class SimulationManager:
|
||||||
if platform is None:
|
if platform is None:
|
||||||
platform = state.get_default_platform()
|
platform = state.get_default_platform()
|
||||||
|
|
||||||
|
if platform not in {"twitter", "reddit"}:
|
||||||
|
raise ValueError(f"不支持的平台: {platform}")
|
||||||
|
|
||||||
sim_dir = self._get_simulation_dir(simulation_id)
|
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):
|
if not os.path.exists(profile_path):
|
||||||
return []
|
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:
|
with open(profile_path, 'r', encoding='utf-8') as f:
|
||||||
return json.load(f)
|
return json.load(f)
|
||||||
|
|
||||||
|
|
|
||||||
|
|
@ -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")
|
||||||
Loading…
Reference in New Issue