search.py 7.9 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245
  1. """Search router - L0 pre-screening and feasibility-first adaptive search (P3-M2)."""
  2. from typing import List, Dict, Any, Optional
  3. from fastapi import APIRouter, HTTPException
  4. from pydantic import BaseModel, Field
  5. from ..services.l0_prescreening import L0PreScreeningEngine
  6. from ..services.feasibility_search import FeasibilityFirstSearch, ParameterRange
  7. router = APIRouter(prefix="/api/search", tags=["Search"])
  8. # Global search instances (in-memory, P3-M2 prototype)
  9. _search_instances: Dict[str, FeasibilityFirstSearch] = {}
  10. _l0_engine = L0PreScreeningEngine()
  11. class L0CheckRequest(BaseModel):
  12. """Request for L0 feasibility check."""
  13. params: Dict[str, Any] = Field(..., description="Parameter values to check")
  14. class L0CheckResponse(BaseModel):
  15. """Response for L0 feasibility check."""
  16. feasible: bool
  17. total_checks: int
  18. passed_checks: int
  19. failed_checks: int
  20. pass_rate: float
  21. results: List[Dict[str, Any]]
  22. risk_items: List[str]
  23. class SearchCreateRequest(BaseModel):
  24. """Request to create a new adaptive search."""
  25. parameters: List[Dict[str, Any]] = Field(..., description="Parameter range definitions")
  26. total_budget: int = Field(default=80, ge=10, le=500)
  27. batch_size: int = Field(default=4, ge=1, le=16)
  28. initial_samples: int = Field(default=16, ge=4, le=100)
  29. objective_metric: str = Field(default="tavg_nm")
  30. objective_direction: str = Field(default="maximize")
  31. seed: int = Field(default=42)
  32. class SearchPointResponse(BaseModel):
  33. """Response for a search point."""
  34. id: int
  35. params: Dict[str, float]
  36. status: str
  37. metrics: Dict[str, float] = Field(default_factory=dict)
  38. batch_id: int
  39. feasible: Optional[bool] = None
  40. class SearchStateResponse(BaseModel):
  41. """Response for search state summary."""
  42. run_id: str
  43. search_method: str
  44. convergence_status: str
  45. total_budget: int
  46. used_budget: int
  47. remaining_budget: int
  48. current_batch: int
  49. total_points: int
  50. pending_points: int
  51. completed_points: int
  52. feasible_points: int
  53. trust_region_active: bool
  54. trust_region_radius: float
  55. best_objective_value: Optional[float] = None
  56. best_point_params: Optional[Dict[str, float]] = None
  57. points_history: List[Dict[str, Any]] = Field(default_factory=list)
  58. infeasible_points: int = 0
  59. failed_points: int = 0
  60. batch_summary: List[Dict[str, Any]] = Field(default_factory=list)
  61. l0_summary: Optional[Dict[str, Any]] = None
  62. class SearchResultRequest(BaseModel):
  63. """Request to report simulation result."""
  64. point_id: int
  65. metrics: Dict[str, float]
  66. status: str = Field(default="ok", description="ok / failed")
  67. @router.post("/l0/check", response_model=L0CheckResponse)
  68. def l0_check(request: L0CheckRequest):
  69. """Run L0 analytic pre-screening on a parameter set."""
  70. report = _l0_engine.evaluate(request.params)
  71. data = report.to_dict()
  72. return L0CheckResponse(**data)
  73. @router.post("/l0/filter")
  74. def l0_filter(param_sets: List[Dict[str, Any]]):
  75. """Filter a list of parameter sets into feasible and infeasible."""
  76. feasible, infeasible = _l0_engine.filter_feasible(param_sets)
  77. return {
  78. "total": len(param_sets),
  79. "feasible_count": len(feasible),
  80. "infeasible_count": len(infeasible),
  81. "feasible": feasible,
  82. "infeasible": infeasible,
  83. }
  84. @router.post("/create")
  85. def create_search(request: SearchCreateRequest):
  86. """Create a new feasibility-first adaptive search run."""
  87. try:
  88. parameters = [
  89. ParameterRange(
  90. name=p["name"],
  91. min_value=float(p["min_value"]),
  92. max_value=float(p["max_value"]),
  93. step=p.get("step"),
  94. unit=p.get("unit", ""),
  95. description=p.get("description", ""),
  96. )
  97. for p in request.parameters
  98. ]
  99. except (KeyError, ValueError) as e:
  100. raise HTTPException(status_code=400, detail=f"Invalid parameter definition: {e}")
  101. search = FeasibilityFirstSearch(
  102. parameters=parameters,
  103. l0_engine=_l0_engine,
  104. total_budget=request.total_budget,
  105. batch_size=request.batch_size,
  106. initial_samples=request.initial_samples,
  107. objective_metric=request.objective_metric,
  108. objective_direction=request.objective_direction,
  109. seed=request.seed,
  110. )
  111. # Generate initial batch
  112. initial_points = search.generate_initial_batch()
  113. _search_instances[search.state.run_id] = search
  114. return {
  115. "run_id": search.state.run_id,
  116. "initial_points": [
  117. {"id": p.id, "params": p.params, "status": p.status, "batch_id": p.batch_id}
  118. for p in initial_points
  119. ],
  120. "state": search.get_state_summary(),
  121. }
  122. @router.get("/{run_id}/state", response_model=SearchStateResponse)
  123. def get_search_state(run_id: str):
  124. """Get current state of a search run."""
  125. if run_id not in _search_instances:
  126. raise HTTPException(status_code=404, detail=f"Search run {run_id} not found")
  127. search = _search_instances[run_id]
  128. return SearchStateResponse(**search.get_state_summary())
  129. @router.post("/{run_id}/next-batch")
  130. def next_batch(run_id: str):
  131. """Select and return the next batch of points to simulate."""
  132. if run_id not in _search_instances:
  133. raise HTTPException(status_code=404, detail=f"Search run {run_id} not found")
  134. search = _search_instances[run_id]
  135. points = search.select_next_batch()
  136. if not points:
  137. return {
  138. "run_id": run_id,
  139. "batch": [],
  140. "convergence_status": search.state.convergence_status,
  141. "message": "No more points to select (budget exhausted or converged)",
  142. }
  143. return {
  144. "run_id": run_id,
  145. "batch_id": search.state.current_batch - 1,
  146. "points": [
  147. {"id": p.id, "params": p.params, "status": p.status, "batch_id": p.batch_id}
  148. for p in points
  149. ],
  150. "state": search.get_state_summary(),
  151. }
  152. @router.post("/{run_id}/report")
  153. def report_result(run_id: str, request: SearchResultRequest):
  154. """Report simulation result for a point."""
  155. if run_id not in _search_instances:
  156. raise HTTPException(status_code=404, detail=f"Search run {run_id} not found")
  157. search = _search_instances[run_id]
  158. search.report_result(request.point_id, request.metrics, request.status)
  159. return {
  160. "run_id": run_id,
  161. "point_id": request.point_id,
  162. "status": "reported",
  163. "state": search.get_state_summary(),
  164. }
  165. @router.get("/{run_id}/points")
  166. def get_points(run_id: str, status: Optional[str] = None):
  167. """Get all points in a search run, optionally filtered by status."""
  168. if run_id not in _search_instances:
  169. raise HTTPException(status_code=404, detail=f"Search run {run_id} not found")
  170. search = _search_instances[run_id]
  171. points = search.state.points
  172. if status:
  173. points = [p for p in points if p.status == status]
  174. return {
  175. "run_id": run_id,
  176. "total": len(points),
  177. "points": [
  178. {
  179. "id": p.id,
  180. "params": p.params,
  181. "status": p.status,
  182. "metrics": p.metrics,
  183. "batch_id": p.batch_id,
  184. "feasible": p.feasibility_report.get("feasible") if p.feasibility_report else None,
  185. }
  186. for p in points
  187. ],
  188. }
  189. @router.get("/{run_id}/export")
  190. def export_search(run_id: str):
  191. """Export full search state for checkpointing."""
  192. if run_id not in _search_instances:
  193. raise HTTPException(status_code=404, detail=f"Search run {run_id} not found")
  194. search = _search_instances[run_id]
  195. return search.export_state()
  196. @router.get("/runs")
  197. def list_runs():
  198. """List all active search runs."""
  199. return {
  200. "runs": [
  201. {"run_id": run_id, "state": search.get_state_summary()}
  202. for run_id, search in _search_instances.items()
  203. ],
  204. }