test_growth_category_tree.py 6.2 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205
  1. from collections.abc import Generator
  2. from contextlib import contextmanager
  3. from types import SimpleNamespace
  4. from unittest.mock import MagicMock
  5. from sqlalchemy import create_engine
  6. from sqlalchemy.orm import Session, sessionmaker
  7. from api.auth_middleware import normal_user_can_access
  8. from api.services import aigc_article_html
  9. from api.services import growth_category_tree as service
  10. from api.services.growth_category_tree import _build_growth_tree
  11. from supply_infra.db.base import Base
  12. from supply_infra.db.models.global_category_content_weight import (
  13. GlobalCategoryContentWeight,
  14. )
  15. from supply_infra.db.models.global_v2 import GlobalCategoryV2
  16. def test_growth_tree_api_is_available_to_normal_users() -> None:
  17. assert normal_user_can_access("GET", "/api/growth-category-tree")
  18. assert normal_user_can_access("GET", "/api/growth-category-tree/38/channel-contents")
  19. assert normal_user_can_access("GET", "/api/growth-channel-content/content-1/view")
  20. def test_build_growth_tree_uses_stable_ids_and_daily_scores() -> None:
  21. categories = [
  22. SimpleNamespace(
  23. stable_id=10,
  24. parent_stable_id=None,
  25. level=1,
  26. name="根分类",
  27. description="根描述",
  28. ),
  29. SimpleNamespace(
  30. stable_id=11,
  31. parent_stable_id=10,
  32. level=2,
  33. name="子分类",
  34. description=None,
  35. ),
  36. SimpleNamespace(
  37. stable_id=12,
  38. parent_stable_id=999,
  39. level=2,
  40. name="孤立分类",
  41. description=None,
  42. ),
  43. ]
  44. weights = [
  45. SimpleNamespace(
  46. stable_id=10,
  47. source_element_count=8,
  48. read_rate_score=0.3,
  49. avg_read_rate_score=1.2,
  50. like_rate_score=0.04,
  51. account_uid_count=2,
  52. channel_content_id_count=5,
  53. cal_fans_num_sum=123.45,
  54. ),
  55. SimpleNamespace(
  56. stable_id=11,
  57. source_element_count=0,
  58. read_rate_score=0.0,
  59. avg_read_rate_score=0.0,
  60. like_rate_score=0.0,
  61. ),
  62. ]
  63. nodes = _build_growth_tree(categories, weights)
  64. assert [node["id"] for node in nodes] == [10, 12]
  65. root = nodes[0]
  66. assert root["weights"] == {
  67. "read_rate": 0.3,
  68. "avg_read_rate": 1.2,
  69. "like_rate": 0.04,
  70. }
  71. assert root["counts"] == {
  72. "read_rate": 8,
  73. "avg_read_rate": 8,
  74. "like_rate": 8,
  75. }
  76. assert root["account_uid_count"] == 2
  77. assert root["channel_content_id_count"] == 5
  78. assert root["cal_fans_num_sum"] == 123.45
  79. assert root["children"][0]["id"] == 11
  80. assert root["children"][0]["counts"]["read_rate"] == 0
  81. assert "hung_word_count" not in root
  82. def test_build_growth_category_tree_serializes_before_session_closes(monkeypatch) -> None:
  83. engine = create_engine("sqlite+pysqlite:///:memory:")
  84. Base.metadata.create_all(
  85. engine,
  86. tables=[
  87. GlobalCategoryV2.__table__,
  88. GlobalCategoryContentWeight.__table__,
  89. ],
  90. )
  91. factory = sessionmaker(bind=engine, expire_on_commit=True)
  92. with factory() as session:
  93. session.add(
  94. GlobalCategoryV2(
  95. id=1,
  96. stable_id=10,
  97. name="根分类",
  98. description=None,
  99. source_type="实质",
  100. level=1,
  101. parent_stable_id=None,
  102. )
  103. )
  104. session.add(
  105. GlobalCategoryContentWeight(
  106. id=1,
  107. stable_id=10,
  108. parent_stable_id=None,
  109. level=1,
  110. biz_dt="20260819",
  111. status="completed",
  112. source_element_count=5,
  113. read_rate_score=0.2,
  114. avg_read_rate_score=1.5,
  115. like_rate_score=0.03,
  116. )
  117. )
  118. session.commit()
  119. @contextmanager
  120. def get_test_session() -> Generator[Session, None, None]:
  121. session = factory()
  122. try:
  123. yield session
  124. session.commit()
  125. finally:
  126. session.close()
  127. monkeypatch.setattr(service, "get_session", get_test_session)
  128. payload = service.build_growth_category_tree()
  129. assert payload["biz_dt"] == "20260819"
  130. assert payload["nodes"][0]["weights"]["avg_read_rate"] == 1.5
  131. def test_list_growth_category_channel_contents_serializes_top_ten(monkeypatch) -> None:
  132. @contextmanager
  133. def get_test_session():
  134. yield MagicMock()
  135. weight_repository = MagicMock()
  136. weight_repository.get_latest_completed_biz_dt.return_value = "20260819"
  137. rank_repository = MagicMock()
  138. rank_repository.list_top.return_value = [
  139. SimpleNamespace(
  140. rank_no=1,
  141. channel_content_id="content-1",
  142. source_element_id=101,
  143. source_category_stable_id=12,
  144. contribution=0.8,
  145. metric_value=0.4,
  146. weighted_score=1.0,
  147. weighted_share=0.25,
  148. )
  149. ]
  150. monkeypatch.setattr(service, "get_session", get_test_session)
  151. monkeypatch.setattr(
  152. service,
  153. "GlobalCategoryContentWeightRepository",
  154. lambda _session: weight_repository,
  155. )
  156. monkeypatch.setattr(
  157. service,
  158. "GlobalCategoryChannelContentRankRepository",
  159. lambda _session: rank_repository,
  160. )
  161. payload = service.list_growth_category_channel_contents(
  162. stable_id=38,
  163. metric="read_rate",
  164. )
  165. assert payload["biz_dt"] == "20260819"
  166. assert payload["total"] == 1
  167. assert payload["items"][0]["weighted_score"] == 1.0
  168. rank_repository.list_top.assert_called_once_with(
  169. biz_dt="20260819",
  170. stable_id=38,
  171. metric_type="read_rate",
  172. )
  173. def test_latest_article_html_uses_parameterized_content_id(monkeypatch) -> None:
  174. connection = MagicMock()
  175. connection.scalar.return_value = b"<html>article</html>"
  176. engine = MagicMock()
  177. engine.connect.return_value.__enter__.return_value = connection
  178. monkeypatch.setattr(aigc_article_html, "_get_aigc_readonly_engine", lambda: engine)
  179. html = aigc_article_html.get_latest_article_html("content-1")
  180. assert html == "<html>article</html>"
  181. parameters = connection.scalar.call_args.args[1]
  182. assert parameters == {"channel_content_id": "content-1"}