graph.py 4.5 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127
  1. from __future__ import annotations
  2. from dataclasses import dataclass
  3. from typing import Any
  4. from langgraph.graph import END, START, StateGraph
  5. from content_agent.business_modules import (
  6. candidate_evidence,
  7. learning_review,
  8. platform_access,
  9. policy_version,
  10. result_source_lookup,
  11. rule_judgment,
  12. run_record,
  13. search_intent,
  14. source_seed,
  15. walk_strategy,
  16. )
  17. from content_agent.interfaces import PlatformSearchClient, PolicyBundleStore, RuntimeFileStore
  18. from content_agent.models import RunState
  19. @dataclass(frozen=True)
  20. class RunDependencies:
  21. runtime: RuntimeFileStore
  22. platform_client: PlatformSearchClient
  23. policy_store: PolicyBundleStore
  24. def build_run_graph(deps: RunDependencies):
  25. graph = StateGraph(RunState)
  26. def load_source(state: RunState) -> dict[str, Any]:
  27. result = source_seed.run(state["trace_id"], state.get("source"), deps.runtime)
  28. return {**result, "current_step": "load_source"}
  29. def plan_queries(state: RunState) -> dict[str, Any]:
  30. queries = search_intent.run(state["trace_id"], state["pattern_seed_pack"], deps.runtime)
  31. return {"queries": queries, "current_step": "plan_queries"}
  32. def search_platform(state: RunState) -> dict[str, Any]:
  33. results = platform_access.run(state["queries"], deps.platform_client)
  34. return {"platform_results": results, "current_step": "search_platform"}
  35. def build_candidates(state: RunState) -> dict[str, Any]:
  36. result = candidate_evidence.run(
  37. state["trace_id"],
  38. state["platform_results"],
  39. state["source_context"],
  40. deps.runtime,
  41. )
  42. return {**result, "current_step": "build_candidates"}
  43. def load_policy(state: RunState) -> dict[str, Any]:
  44. bundle = policy_version.run(state["policy_bundle_version"], deps.policy_store)
  45. return {"policy_bundle": bundle, "current_step": "load_policy"}
  46. def evaluate_rules(state: RunState) -> dict[str, Any]:
  47. decisions = rule_judgment.run(
  48. state["trace_id"],
  49. state["evidence_bundles"],
  50. state["policy_bundle"],
  51. deps.runtime,
  52. )
  53. return {"rule_decisions": decisions, "current_step": "evaluate_rules"}
  54. def plan_walk(state: RunState) -> dict[str, Any]:
  55. result = walk_strategy.run(
  56. state["pattern_seed_pack"],
  57. state["queries"],
  58. state["candidates"],
  59. state["rule_decisions"],
  60. )
  61. return {**result, "current_step": "plan_walk"}
  62. def record_run(state: RunState) -> dict[str, Any]:
  63. result = run_record.run(
  64. state["trace_id"],
  65. state["queries"],
  66. state["candidates"],
  67. state["rule_decisions"],
  68. state["source_edge_basis"],
  69. deps.runtime,
  70. )
  71. return {**result, "current_step": "record_run"}
  72. def commit_results(state: RunState) -> dict[str, Any]:
  73. final_output = result_source_lookup.run(
  74. state["trace_id"],
  75. state["candidates"],
  76. state["media_assets"],
  77. state["rule_decisions"],
  78. state["source_edges"],
  79. state["search_clues"],
  80. deps.runtime,
  81. )
  82. return {"final_output": final_output, "current_step": "commit_results"}
  83. def review_strategy(state: RunState) -> dict[str, Any]:
  84. review = learning_review.run(state["trace_id"], deps.runtime)
  85. return {"strategy_review": review, "current_step": "review_strategy", "status": "success"}
  86. graph.add_node("load_source", load_source)
  87. graph.add_node("plan_queries", plan_queries)
  88. graph.add_node("search_platform", search_platform)
  89. graph.add_node("build_candidates", build_candidates)
  90. graph.add_node("load_policy", load_policy)
  91. graph.add_node("evaluate_rules", evaluate_rules)
  92. graph.add_node("plan_walk", plan_walk)
  93. graph.add_node("record_run", record_run)
  94. graph.add_node("commit_results", commit_results)
  95. graph.add_node("review_strategy", review_strategy)
  96. graph.add_edge(START, "load_source")
  97. graph.add_edge("load_source", "plan_queries")
  98. graph.add_edge("plan_queries", "search_platform")
  99. graph.add_edge("search_platform", "build_candidates")
  100. graph.add_edge("build_candidates", "load_policy")
  101. graph.add_edge("load_policy", "evaluate_rules")
  102. graph.add_edge("evaluate_rules", "plan_walk")
  103. graph.add_edge("plan_walk", "record_run")
  104. graph.add_edge("record_run", "commit_results")
  105. graph.add_edge("commit_results", "review_strategy")
  106. graph.add_edge("review_strategy", END)
  107. return graph.compile()