generation.py 3.9 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103
  1. """Plan generation API router (rule engine + parameter registry)."""
  2. from fastapi import APIRouter, Depends, HTTPException
  3. from sqlalchemy.orm import Session
  4. from ..database import get_db
  5. from ..models.project import Project
  6. from ..schemas.generation import (
  7. ParameterRegistryResponse,
  8. RangeRecommendRequest, RangeRecommendation,
  9. PlanGenerateRequest, PlanGenerateResponse,
  10. )
  11. from ..services.rule_engine import (
  12. get_parameter_registry, get_parameter,
  13. recommend_range, generate_plan, BoundaryConditions,
  14. )
  15. router = APIRouter(prefix="/api", tags=["generation"])
  16. # ---------------------------------------------------------------------------
  17. # Parameter registry
  18. # ---------------------------------------------------------------------------
  19. @router.get("/scan-parameters", response_model=ParameterRegistryResponse)
  20. def list_scan_parameters(category: str | None = None):
  21. """List all scannable Motor-CAD parameters with their metadata.
  22. Used by the frontend plan editor to populate variable dropdowns.
  23. """
  24. params = get_parameter_registry()
  25. if category:
  26. params = [p for p in params if p["category"].lower() == category.lower()]
  27. return ParameterRegistryResponse(total=len(params), parameters=params)
  28. # ---------------------------------------------------------------------------
  29. # Range recommendation
  30. # ---------------------------------------------------------------------------
  31. @router.post("/recommend-range", response_model=RangeRecommendation)
  32. def recommend_scan_range(request: RangeRecommendRequest):
  33. """Recommend a scan range for a parameter based on boundary conditions."""
  34. p = get_parameter(request.parameter_name)
  35. if p is None:
  36. raise HTTPException(
  37. status_code=404,
  38. detail=f"Unknown parameter: {request.parameter_name}",
  39. )
  40. bc = BoundaryConditions.from_dict(request.boundary_conditions or {})
  41. rng = recommend_range(request.parameter_name, bc)
  42. return RangeRecommendation(**rng)
  43. # ---------------------------------------------------------------------------
  44. # Plan generation
  45. # ---------------------------------------------------------------------------
  46. @router.post("/generate-plan", response_model=PlanGenerateResponse)
  47. def generate_plan_from_bc(request: PlanGenerateRequest):
  48. """Generate a recommended simulation plan from boundary conditions.
  49. Uses the rule engine to recommend scan ranges for each parameter and
  50. returns a plan draft that can be edited and then saved via the plans API.
  51. """
  52. bc = BoundaryConditions.from_dict(request.boundary_conditions or {})
  53. plan = generate_plan(
  54. bc=bc,
  55. param_names=request.parameter_names,
  56. model_path=request.model_path,
  57. )
  58. return PlanGenerateResponse(**plan)
  59. @router.post("/projects/{project_id}/generate-plan", response_model=PlanGenerateResponse)
  60. def generate_plan_for_project(
  61. project_id: int,
  62. request: PlanGenerateRequest,
  63. db: Session = Depends(get_db),
  64. ):
  65. """Generate a plan for an existing project using its boundary conditions.
  66. The project's stored boundary conditions are used as the base; any
  67. boundary_conditions in the request override them.
  68. """
  69. project = db.query(Project).filter(Project.id == project_id).first()
  70. if not project:
  71. raise HTTPException(status_code=404, detail="Project not found")
  72. # Merge: project BC as base, request BC overrides
  73. merged_bc = dict(project.get_boundary_conditions())
  74. if request.boundary_conditions:
  75. merged_bc.update(request.boundary_conditions)
  76. # Ensure topology from project
  77. merged_bc.setdefault("topology", project.topology or "SSSR")
  78. bc = BoundaryConditions.from_dict(merged_bc)
  79. model_path = request.model_path or project.model_path or ""
  80. plan = generate_plan(
  81. bc=bc,
  82. param_names=request.parameter_names,
  83. model_path=model_path,
  84. )
  85. return PlanGenerateResponse(**plan)