165 lines
4.9 KiB
Go
165 lines
4.9 KiB
Go
package api
|
|
|
|
import (
|
|
"encoding/json"
|
|
"net/http"
|
|
"strconv"
|
|
|
|
"github.com/gorilla/mux"
|
|
"go.temporal.io/sdk/client"
|
|
"github.com/rockliang/poimen/workflows/pkg/db"
|
|
"github.com/rockliang/poimen/workflows/statemachine"
|
|
)
|
|
|
|
// QueryWorkflowGraphRequest matches frontend payload
|
|
type QueryWorkflowGraphRequest struct {
|
|
Query string `json:"query"`
|
|
SearchType string `json:"search_type"`
|
|
RelationType string `json:"relation_type"`
|
|
Version int `json:"version"`
|
|
ConfidenceFloor float64 `json:"confidence_floor"`
|
|
TopK int `json:"top_k"`
|
|
FindPaths bool `json:"find_paths"`
|
|
TargetNodeID string `json:"target_node_id"`
|
|
MaxPathDepth int `json:"max_path_depth"`
|
|
RankingProfile string `json:"ranking_profile"`
|
|
IncludeReasoning bool `json:"include_reasoning"`
|
|
}
|
|
|
|
// QueryWorkflowGraphHandler handles POST /workflows/{id}/query
|
|
func QueryWorkflowGraphHandler(temporalClient client.Client, dbClient *db.Client) http.HandlerFunc {
|
|
return func(w http.ResponseWriter, r *http.Request) {
|
|
if r.Method != http.MethodPost {
|
|
http.Error(w, "Method not allowed", http.StatusMethodNotAllowed)
|
|
return
|
|
}
|
|
|
|
workflowID := mux.Vars(r)["id"]
|
|
if workflowID == "" {
|
|
http.Error(w, "Missing workflow ID", http.StatusBadRequest)
|
|
return
|
|
}
|
|
|
|
var req QueryWorkflowGraphRequest
|
|
if err := json.NewDecoder(r.Body).Decode(&req); err != nil {
|
|
http.Error(w, "Invalid request body", http.StatusBadRequest)
|
|
return
|
|
}
|
|
|
|
// Defaults
|
|
if req.SearchType == "" {
|
|
req.SearchType = "edges"
|
|
}
|
|
if req.ConfidenceFloor == 0 {
|
|
req.ConfidenceFloor = 0.5
|
|
}
|
|
if req.TopK == 0 {
|
|
req.TopK = 10
|
|
}
|
|
if req.MaxPathDepth == 0 {
|
|
req.MaxPathDepth = 3
|
|
}
|
|
if req.RankingProfile == "" {
|
|
req.RankingProfile = "default"
|
|
}
|
|
|
|
// Get latest version if not specified
|
|
if req.Version == 0 {
|
|
workflow, err := dbClient.GetWorkflow(r.Context(), workflowID)
|
|
if err != nil {
|
|
http.Error(w, "Failed to get workflow", http.StatusInternalServerError)
|
|
return
|
|
}
|
|
req.Version = workflow.Version
|
|
}
|
|
|
|
// Start Temporal workflow
|
|
workflowInput := statemachine.WorkflowGraphQueryInput{
|
|
WorkflowID: workflowID,
|
|
Query: req.Query,
|
|
SearchType: req.SearchType,
|
|
RelationType: req.RelationType,
|
|
Version: req.Version,
|
|
ConfidenceFloor: req.ConfidenceFloor,
|
|
TopK: req.TopK,
|
|
FindPaths: req.FindPaths,
|
|
TargetNodeID: req.TargetNodeID,
|
|
MaxPathDepth: req.MaxPathDepth,
|
|
RankingProfile: req.RankingProfile,
|
|
IncludeReasoning: req.IncludeReasoning,
|
|
}
|
|
|
|
options := client.StartWorkflowOptions{
|
|
ID: "graph-query-" + workflowID + "-" + strconv.FormatInt(int64(req.Version), 10),
|
|
TaskQueue: "poimen-default",
|
|
}
|
|
|
|
// Execute workflow (blocking)
|
|
run, err := temporalClient.ExecuteWorkflow(
|
|
r.Context(),
|
|
options,
|
|
statemachine.WorkflowGraphQuery,
|
|
workflowInput,
|
|
)
|
|
if err != nil {
|
|
http.Error(w, "Failed to start workflow", http.StatusInternalServerError)
|
|
return
|
|
}
|
|
|
|
var output statemachine.WorkflowGraphQueryOutput
|
|
if err := run.Get(r.Context(), &output); err != nil {
|
|
http.Error(w, "Workflow execution failed", http.StatusInternalServerError)
|
|
return
|
|
}
|
|
|
|
w.Header().Set("Content-Type", "application/json")
|
|
json.NewEncoder(w).Encode(output)
|
|
}
|
|
}
|
|
|
|
// GetWorkflowRelationVersionsHandler handles GET /workflows/{id}/relations/{edge_id}/versions
|
|
func GetWorkflowRelationVersionsHandler(dbClient *db.Client) http.HandlerFunc {
|
|
return func(w http.ResponseWriter, r *http.Request) {
|
|
if r.Method != http.MethodGet {
|
|
http.Error(w, "Method not allowed", http.StatusMethodNotAllowed)
|
|
return
|
|
}
|
|
|
|
workflowID := mux.Vars(r)["id"]
|
|
edgeID := mux.Vars(r)["edge_id"]
|
|
|
|
if workflowID == "" || edgeID == "" {
|
|
http.Error(w, "Missing parameters", http.StatusBadRequest)
|
|
return
|
|
}
|
|
|
|
// Query relation versions from DB
|
|
versions, err := dbClient.GetRelationVersions(r.Context(), workflowID, edgeID)
|
|
if err != nil {
|
|
http.Error(w, "Failed to fetch versions", http.StatusInternalServerError)
|
|
return
|
|
}
|
|
|
|
w.Header().Set("Content-Type", "application/json")
|
|
json.NewEncoder(w).Encode(map[string]interface{}{
|
|
"workflow_id": workflowID,
|
|
"edge_id": edgeID,
|
|
"versions": versions,
|
|
"total_count": len(versions),
|
|
})
|
|
}
|
|
}
|
|
|
|
// RegisterGraphQueryHandlers registers all graph query endpoints
|
|
func RegisterGraphQueryHandlers(
|
|
router *mux.Router,
|
|
temporalClient client.Client,
|
|
dbClient *db.Client,
|
|
) {
|
|
// Query workflow relations via GraphRAG
|
|
router.HandleFunc("/workflows/{id}/query", QueryWorkflowGraphHandler(temporalClient, dbClient)).Methods("POST")
|
|
|
|
// Get relation version history
|
|
router.HandleFunc("/workflows/{id}/relations/{edge_id}/versions", GetWorkflowRelationVersionsHandler(dbClient)).Methods("GET")
|
|
}
|