test_offline_regression.py 4.5 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111
  1. from __future__ import annotations
  2. import sys
  3. import unittest
  4. from pathlib import Path
  5. import pandas as pd
  6. PROJECT_DIR = Path(__file__).resolve().parents[1]
  7. WORKSPACE_ROOT = PROJECT_DIR.parent
  8. for path in (PROJECT_DIR, WORKSPACE_ROOT):
  9. if str(path) not in sys.path:
  10. sys.path.insert(0, str(path))
  11. from scripts.build_report import group_summary
  12. from src.social_logic import aggregate_groups, aggregate_relationships
  13. from src.user_input import load_users
  14. FIXTURE_DIR = PROJECT_DIR / "fixtures" / "v2_220"
  15. SAMPLE_PATH = PROJECT_DIR / "samples" / "v2_220_20260810.csv"
  16. class OfflineRegressionTest(unittest.TestCase):
  17. @classmethod
  18. def setUpClass(cls) -> None:
  19. cls.users = load_users(SAMPLE_PATH, lookback_days=180)
  20. cls.behavior = pd.read_csv(FIXTURE_DIR / "user_detail.csv")
  21. cls.relationships = pd.read_csv(FIXTURE_DIR / "relationships.csv")
  22. cls.groups = pd.read_csv(FIXTURE_DIR / "social_group_features.csv")
  23. cls.expected_social = pd.read_csv(FIXTURE_DIR / "social_user_features.csv")
  24. def test_sample_contract(self) -> None:
  25. self.assertEqual(len(self.users), 220)
  26. self.assertEqual(self.users["mid"].nunique(), 220)
  27. self.assertTrue((self.users["window_end_dt"] < self.users["anchor_dt"]).all())
  28. counts = self.users.groupby(["user_type", "sub_category"]).size().to_dict()
  29. self.assertEqual(counts[("明确举报人", "明确举报人")], 7)
  30. self.assertEqual(counts[("审核人", "审核人")], 13)
  31. self.assertEqual(counts[("正常用户", "正常用户")], 100)
  32. suspected = sum(value for (group, _), value in counts.items() if group == "疑似举报人")
  33. self.assertEqual(suspected, 100)
  34. def test_long_retention_play_and_share_sources(self) -> None:
  35. builder = (PROJECT_DIR / "src" / "feature_sql.py").read_text(encoding="utf-8")
  36. dataworks = (
  37. PROJECT_DIR / "sql" / "dataworks_single_user_180d_all_in_one.sql"
  38. ).read_text(encoding="utf-8")
  39. for sql in (builder, dataworks):
  40. self.assertIn("loghubods.ods_video_play_log_day", sql)
  41. self.assertIn("s.topic = 'share'", sql)
  42. self.assertNotIn("LEFT JOIN loghubods.video_play_log p", sql)
  43. def test_social_aggregation_matches_saved_result(self) -> None:
  44. relationships = aggregate_relationships(self.users, self.relationships)
  45. groups = aggregate_groups(self.users, self.groups)
  46. actual = relationships.merge(groups, on="用户id", validate="one_to_one")
  47. expected = self.expected_social
  48. self.assertEqual(len(actual), 220)
  49. self.assertEqual(actual["用户id"].nunique(), 220)
  50. columns = [
  51. "社交关系点击次数",
  52. "风险来源点击次数",
  53. "同一卡片最大连续点击次数",
  54. "有效来源群数",
  55. "来源群访问UV合计",
  56. "来源群有效播放UV合计",
  57. ]
  58. merged = actual[["用户id", *columns]].merge(
  59. expected[["用户id", *columns]],
  60. on="用户id",
  61. suffixes=("_actual", "_expected"),
  62. validate="one_to_one",
  63. )
  64. for column in columns:
  65. pd.testing.assert_series_equal(
  66. pd.to_numeric(merged[f"{column}_actual"], errors="coerce").fillna(0),
  67. pd.to_numeric(merged[f"{column}_expected"], errors="coerce").fillna(0),
  68. check_names=False,
  69. )
  70. def test_population_group_facts(self) -> None:
  71. detail = self.behavior.merge(
  72. self.expected_social,
  73. on="用户id",
  74. how="left",
  75. validate="one_to_one",
  76. )
  77. summary = group_summary(detail).set_index("用户分类")
  78. expected = {
  79. "明确举报人": (7, 8, 149, 0),
  80. "审核人": (13, 2, 54, 19),
  81. "疑似举报人": (100, 316, 5560, 2543),
  82. "正常用户": (100, 310, 6404, 2760),
  83. }
  84. for group, values in expected.items():
  85. self.assertEqual(int(summary.at[group, "总人数"]), values[0])
  86. self.assertEqual(int(summary.at[group, "访问来源群数合计"]), values[1])
  87. self.assertEqual(
  88. int(summary.at[group, "来源群访问用户数合计(按群UV求和)"]),
  89. values[2],
  90. )
  91. self.assertEqual(
  92. int(summary.at[group, "来源群有效播放用户数合计(按群UV求和)"]),
  93. values[3],
  94. )
  95. if __name__ == "__main__":
  96. unittest.main()