485 lines
13 KiB
Go
485 lines
13 KiB
Go
// Package proxy provides request routing and forwarding.
|
|||
|
|
package proxy
|
||
|
|
|
||
|
|
import (
|
||
|
|
"bytes"
|
||
|
|
"context"
|
||
|
|
"encoding/json"
|
||
|
|
"fmt"
|
||
|
|
"io"
|
||
|
|
"net/http"
|
||
|
|
"time"
|
||
|
|
)
|
||
|
|
|
||
|
|
// WorkflowRequest represents a workflow execution request
|
||
|
|
type WorkflowRequest struct {
|
||
|
|
// Workflow ID or name
|
||
|
|
Workflow string `json:"workflow"`
|
||
|
|
|
||
|
|
// Input parameters for the workflow
|
||
|
|
Input map[string]interface{} `json:"input"`
|
||
|
|
|
||
|
|
// Optional: timeout in seconds
|
||
|
|
Timeout int `json:"timeout,omitempty"`
|
||
|
|
|
||
|
|
// Optional: wait for result (default: true)
|
||
|
|
Wait *bool `json:"wait,omitempty"`
|
||
|
|
}
|
||
|
|
|
||
|
|
// WorkflowResponse represents the response from workflow execution
|
||
|
|
type WorkflowResponse struct {
|
||
|
|
// Workflow execution ID
|
||
|
|
ID string `json:"id"`
|
||
|
|
|
||
|
|
// Workflow name
|
||
|
|
Workflow string `json:"workflow"`
|
||
|
|
|
||
|
|
// Execution status: pending, running, completed, failed
|
||
|
|
Status string `json:"status"`
|
||
|
|
|
||
|
|
// Output of the workflow
|
||
|
|
Output interface{} `json:"output,omitempty"`
|
||
|
|
|
||
|
|
// Error message if workflow failed
|
||
|
|
Error string `json:"error,omitempty"`
|
||
|
|
|
||
|
|
// Timestamp when workflow was created
|
||
|
|
CreatedAt time.Time `json:"created_at"`
|
||
|
|
|
||
|
|
// Timestamp when workflow completed
|
||
|
|
CompletedAt *time.Time `json:"completed_at,omitempty"`
|
||
|
|
}
|
||
|
|
|
||
|
|
// PredefinedWorkflow defines a workflow template that combines multiple API calls
|
||
|
|
type PredefinedWorkflow struct {
|
||
|
|
Name string
|
||
|
|
Description string
|
||
|
|
Handler func(*http.Request, *Handler, map[string]interface{}) (interface{}, error)
|
||
|
|
}
|
||
|
|
|
||
|
|
// handleWorkflow handles the /workflows endpoint
|
||
|
|
// It accepts workflow definitions and orchestrates API calls
|
||
|
|
func (h *Handler) handleWorkflow(w http.ResponseWriter, r *http.Request) {
|
||
|
|
// Only POST is supported
|
||
|
|
if r.Method != "POST" {
|
||
|
|
w.Header().Set("Content-Type", "application/problem+json")
|
||
|
|
w.WriteHeader(http.StatusMethodNotAllowed)
|
||
|
|
fmt.Fprintf(w, `{"type":"https://api.example.com/problems/method-not-allowed","title":"Method Not Allowed","status":405,"detail":"Only POST is supported for /workflows"}`)
|
||
|
|
return
|
||
|
|
}
|
||
|
|
|
||
|
|
// Parse request body
|
||
|
|
var workflowReq WorkflowRequest
|
||
|
|
if err := json.NewDecoder(r.Body).Decode(&workflowReq); err != nil {
|
||
|
|
writeProblemDetail(w, http.StatusBadRequest, "https://api.example.com/problems/invalid-workflow-request", "Invalid Workflow Request", "Failed to parse workflow request: "+err.Error(), nil)
|
||
|
|
return
|
||
|
|
}
|
||
|
|
|
||
|
|
// Validate workflow name
|
||
|
|
if workflowReq.Workflow == "" {
|
||
|
|
writeProblemDetail(w, http.StatusBadRequest, "https://api.example.com/problems/missing-workflow", "Missing Workflow", "The 'workflow' field is required", nil)
|
||
|
|
return
|
||
|
|
}
|
||
|
|
|
||
|
|
// Get predefined workflow
|
||
|
|
workflow, ok := h.getWorkflow(workflowReq.Workflow)
|
||
|
|
if !ok {
|
||
|
|
availableWorkflows := h.getAvailableWorkflows()
|
||
|
|
writeProblemDetail(w, http.StatusBadRequest, "https://api.example.com/problems/unknown-workflow", "Unknown Workflow", fmt.Sprintf("Workflow %q is not available", workflowReq.Workflow), availableWorkflows)
|
||
|
|
return
|
||
|
|
}
|
||
|
|
|
||
|
|
// Default wait to true
|
||
|
|
wait := true
|
||
|
|
if workflowReq.Wait != nil {
|
||
|
|
wait = *workflowReq.Wait
|
||
|
|
}
|
||
|
|
|
||
|
|
// Set default timeout if not provided
|
||
|
|
timeout := time.Duration(30) * time.Second
|
||
|
|
if workflowReq.Timeout > 0 {
|
||
|
|
timeout = time.Duration(workflowReq.Timeout) * time.Second
|
||
|
|
}
|
||
|
|
|
||
|
|
// Create a context with timeout for workflow execution
|
||
|
|
ctx, cancel := context.WithTimeout(context.Background(), timeout)
|
||
|
|
defer cancel()
|
||
|
|
|
||
|
|
// Execute workflow
|
||
|
|
output, err := workflow.Handler(r.WithContext(ctx), h, workflowReq.Input)
|
||
|
|
|
||
|
|
// Build response
|
||
|
|
workflowResp := WorkflowResponse{
|
||
|
|
ID: generateWorkflowID(),
|
||
|
|
Workflow: workflowReq.Workflow,
|
||
|
|
CreatedAt: time.Now(),
|
||
|
|
}
|
||
|
|
|
||
|
|
if err != nil {
|
||
|
|
workflowResp.Status = "failed"
|
||
|
|
workflowResp.Error = err.Error()
|
||
|
|
} else {
|
||
|
|
if wait {
|
||
|
|
workflowResp.Status = "completed"
|
||
|
|
workflowResp.Output = output
|
||
|
|
now := time.Now()
|
||
|
|
workflowResp.CompletedAt = &now
|
||
|
|
} else {
|
||
|
|
workflowResp.Status = "pending"
|
||
|
|
}
|
||
|
|
}
|
||
|
|
|
||
|
|
// Write response
|
||
|
|
w.Header().Set("Content-Type", "application/json")
|
||
|
|
if err != nil {
|
||
|
|
w.WriteHeader(http.StatusInternalServerError)
|
||
|
|
} else {
|
||
|
|
w.WriteHeader(http.StatusOK)
|
||
|
|
}
|
||
|
|
json.NewEncoder(w).Encode(workflowResp)
|
||
|
|
}
|
||
|
|
|
||
|
|
// getWorkflow returns a predefined workflow by name
|
||
|
|
func (h *Handler) getWorkflow(name string) (*PredefinedWorkflow, bool) {
|
||
|
|
workflows := h.getPredefinedWorkflows()
|
||
|
|
for _, wf := range workflows {
|
||
|
|
if wf.Name == name {
|
||
|
|
return &wf, true
|
||
|
|
}
|
||
|
|
}
|
||
|
|
return nil, false
|
||
|
|
}
|
||
|
|
|
||
|
|
// getPredefinedWorkflows returns all available workflows
|
||
|
|
func (h *Handler) getPredefinedWorkflows() []PredefinedWorkflow {
|
||
|
|
return []PredefinedWorkflow{
|
||
|
|
{
|
||
|
|
Name: "chat-and-embed",
|
||
|
|
Description: "Chat with a model and then embed the response",
|
||
|
|
Handler: h.chatAndEmbedWorkflow,
|
||
|
|
},
|
||
|
|
{
|
||
|
|
Name: "multi-model-chat",
|
||
|
|
Description: "Chat with multiple models sequentially",
|
||
|
|
Handler: h.multiModelChatWorkflow,
|
||
|
|
},
|
||
|
|
{
|
||
|
|
Name: "rag-pipeline",
|
||
|
|
Description: "RAG pipeline: embed query, rerank, then chat with context",
|
||
|
|
Handler: h.ragPipelineWorkflow,
|
||
|
|
},
|
||
|
|
{
|
||
|
|
Name: "batch-embeddings",
|
||
|
|
Description: "Generate embeddings for multiple texts",
|
||
|
|
Handler: h.batchEmbeddingsWorkflow,
|
||
|
|
},
|
||
|
|
}
|
||
|
|
}
|
||
|
|
|
||
|
|
// getAvailableWorkflows returns a list of available workflow names
|
||
|
|
func (h *Handler) getAvailableWorkflows() []string {
|
||
|
|
workflows := h.getPredefinedWorkflows()
|
||
|
|
names := make([]string, len(workflows))
|
||
|
|
for i, wf := range workflows {
|
||
|
|
names[i] = wf.Name
|
||
|
|
}
|
||
|
|
return names
|
||
|
|
}
|
||
|
|
|
||
|
|
// Workflow implementations
|
||
|
|
|
||
|
|
// chatAndEmbedWorkflow: Chat with a model, then embed the response
|
||
|
|
func (h *Handler) chatAndEmbedWorkflow(r *http.Request, handler *Handler, input map[string]interface{}) (interface{}, error) {
|
||
|
|
model, ok := input["model"].(string)
|
||
|
|
if !ok || model == "" {
|
||
|
|
return nil, fmt.Errorf("missing required parameter: model")
|
||
|
|
}
|
||
|
|
|
||
|
|
embedModel, ok := input["embed_model"].(string)
|
||
|
|
if !ok {
|
||
|
|
embedModel = "nomic-ai/nomic-embed-text-v2-moe"
|
||
|
|
}
|
||
|
|
|
||
|
|
messages, ok := input["messages"].([]interface{})
|
||
|
|
if !ok {
|
||
|
|
return nil, fmt.Errorf("missing required parameter: messages")
|
||
|
|
}
|
||
|
|
|
||
|
|
// Step 1: Chat
|
||
|
|
chatReq := map[string]interface{}{
|
||
|
|
"model": model,
|
||
|
|
"messages": messages,
|
||
|
|
}
|
||
|
|
|
||
|
|
chatBody, _ := json.Marshal(chatReq)
|
||
|
|
chatHTTPReq, _ := http.NewRequest("POST", "/v1/chat/completions", io.NopCloser(bytes.NewReader(chatBody)))
|
||
|
|
chatHTTPReq.Header.Set("Content-Type", "application/json")
|
||
|
|
|
||
|
|
// Create a response writer to capture the chat response
|
||
|
|
chatResp := &responseCapture{}
|
||
|
|
handler.ServeHTTP(chatResp, chatHTTPReq)
|
||
|
|
|
||
|
|
var chatResult map[string]interface{}
|
||
|
|
if err := json.Unmarshal(chatResp.body.Bytes(), &chatResult); err != nil {
|
||
|
|
return nil, fmt.Errorf("failed to parse chat response: %v", err)
|
||
|
|
}
|
||
|
|
|
||
|
|
// Extract message content
|
||
|
|
var messageContent string
|
||
|
|
if choices, ok := chatResult["choices"].([]interface{}); ok && len(choices) > 0 {
|
||
|
|
if choice, ok := choices[0].(map[string]interface{}); ok {
|
||
|
|
if message, ok := choice["message"].(map[string]interface{}); ok {
|
||
|
|
if content, ok := message["content"].(string); ok {
|
||
|
|
messageContent = content
|
||
|
|
}
|
||
|
|
}
|
||
|
|
}
|
||
|
|
}
|
||
|
|
|
||
|
|
// Step 2: Embed the response
|
||
|
|
embedReq := map[string]interface{}{
|
||
|
|
"model": embedModel,
|
||
|
|
"input": messageContent,
|
||
|
|
}
|
||
|
|
|
||
|
|
embedBody, _ := json.Marshal(embedReq)
|
||
|
|
embedHTTPReq, _ := http.NewRequest("POST", "/v1/embeddings", io.NopCloser(bytes.NewReader(embedBody)))
|
||
|
|
embedHTTPReq.Header.Set("Content-Type", "application/json")
|
||
|
|
|
||
|
|
embedResp := &responseCapture{}
|
||
|
|
handler.ServeHTTP(embedResp, embedHTTPReq)
|
||
|
|
|
||
|
|
var embedResult map[string]interface{}
|
||
|
|
if err := json.Unmarshal(embedResp.body.Bytes(), &embedResult); err != nil {
|
||
|
|
return nil, fmt.Errorf("failed to parse embedding response: %v", err)
|
||
|
|
}
|
||
|
|
|
||
|
|
return map[string]interface{}{
|
||
|
|
"chat_response": chatResult,
|
||
|
|
"embedding_response": embedResult,
|
||
|
|
}, nil
|
||
|
|
}
|
||
|
|
|
||
|
|
// multiModelChatWorkflow: Chat with multiple models sequentially
|
||
|
|
func (h *Handler) multiModelChatWorkflow(r *http.Request, handler *Handler, input map[string]interface{}) (interface{}, error) {
|
||
|
|
models, ok := input["models"].([]interface{})
|
||
|
|
if !ok || len(models) == 0 {
|
||
|
|
return nil, fmt.Errorf("missing required parameter: models (array)")
|
||
|
|
}
|
||
|
|
|
||
|
|
messages, ok := input["messages"].([]interface{})
|
||
|
|
if !ok {
|
||
|
|
return nil, fmt.Errorf("missing required parameter: messages")
|
||
|
|
}
|
||
|
|
|
||
|
|
results := make([]map[string]interface{}, 0)
|
||
|
|
|
||
|
|
for _, modelInterface := range models {
|
||
|
|
model, ok := modelInterface.(string)
|
||
|
|
if !ok {
|
||
|
|
continue
|
||
|
|
}
|
||
|
|
|
||
|
|
chatReq := map[string]interface{}{
|
||
|
|
"model": model,
|
||
|
|
"messages": messages,
|
||
|
|
}
|
||
|
|
|
||
|
|
chatBody, _ := json.Marshal(chatReq)
|
||
|
|
chatHTTPReq, _ := http.NewRequest("POST", "/v1/chat/completions", io.NopCloser(bytes.NewReader(chatBody)))
|
||
|
|
chatHTTPReq.Header.Set("Content-Type", "application/json")
|
||
|
|
|
||
|
|
chatResp := &responseCapture{}
|
||
|
|
handler.ServeHTTP(chatResp, chatHTTPReq)
|
||
|
|
|
||
|
|
var chatResult map[string]interface{}
|
||
|
|
if err := json.Unmarshal(chatResp.body.Bytes(), &chatResult); err != nil {
|
||
|
|
results = append(results, map[string]interface{}{
|
||
|
|
"model": model,
|
||
|
|
"error": err.Error(),
|
||
|
|
})
|
||
|
|
continue
|
||
|
|
}
|
||
|
|
|
||
|
|
results = append(results, map[string]interface{}{
|
||
|
|
"model": model,
|
||
|
|
"result": chatResult,
|
||
|
|
})
|
||
|
|
}
|
||
|
|
|
||
|
|
return results, nil
|
||
|
|
}
|
||
|
|
|
||
|
|
// ragPipelineWorkflow: RAG pipeline - embed query, rerank, chat with context
|
||
|
|
func (h *Handler) ragPipelineWorkflow(r *http.Request, handler *Handler, input map[string]interface{}) (interface{}, error) {
|
||
|
|
query, ok := input["query"].(string)
|
||
|
|
if !ok || query == "" {
|
||
|
|
return nil, fmt.Errorf("missing required parameter: query")
|
||
|
|
}
|
||
|
|
|
||
|
|
documents, ok := input["documents"].([]interface{})
|
||
|
|
if !ok {
|
||
|
|
return nil, fmt.Errorf("missing required parameter: documents")
|
||
|
|
}
|
||
|
|
|
||
|
|
model, ok := input["model"].(string)
|
||
|
|
if !ok {
|
||
|
|
model = "reasoning"
|
||
|
|
}
|
||
|
|
|
||
|
|
rerankModel, ok := input["rerank_model"].(string)
|
||
|
|
if !ok {
|
||
|
|
rerankModel = "BAAI/bge-reranker-base"
|
||
|
|
}
|
||
|
|
|
||
|
|
topK := 3
|
||
|
|
if tk, ok := input["top_k"].(float64); ok {
|
||
|
|
topK = int(tk)
|
||
|
|
}
|
||
|
|
|
||
|
|
// Step 1: Rerank documents based on query
|
||
|
|
rerankReq := map[string]interface{}{
|
||
|
|
"model": rerankModel,
|
||
|
|
"query": query,
|
||
|
|
"texts": documents,
|
||
|
|
"top_k": topK,
|
||
|
|
}
|
||
|
|
|
||
|
|
rerankBody, _ := json.Marshal(rerankReq)
|
||
|
|
rerankHTTPReq, _ := http.NewRequest("POST", "/v1/rerank", io.NopCloser(bytes.NewReader(rerankBody)))
|
||
|
|
rerankHTTPReq.Header.Set("Content-Type", "application/json")
|
||
|
|
|
||
|
|
rerankResp := &responseCapture{}
|
||
|
|
handler.ServeHTTP(rerankResp, rerankHTTPReq)
|
||
|
|
|
||
|
|
var rerankResult map[string]interface{}
|
||
|
|
if err := json.Unmarshal(rerankResp.body.Bytes(), &rerankResult); err != nil {
|
||
|
|
return nil, fmt.Errorf("failed to parse rerank response: %v", err)
|
||
|
|
}
|
||
|
|
|
||
|
|
// Extract top documents
|
||
|
|
var topDocs []string
|
||
|
|
if results, ok := rerankResult["results"].([]interface{}); ok {
|
||
|
|
for i, resultInterface := range results {
|
||
|
|
if i >= topK {
|
||
|
|
break
|
||
|
|
}
|
||
|
|
if result, ok := resultInterface.(map[string]interface{}); ok {
|
||
|
|
if text, ok := result["text"].(string); ok {
|
||
|
|
topDocs = append(topDocs, text)
|
||
|
|
}
|
||
|
|
}
|
||
|
|
}
|
||
|
|
}
|
||
|
|
|
||
|
|
// Step 2: Chat with context
|
||
|
|
context := fmt.Sprintf("Context from documents:\n%v\n\nQuery: %s", topDocs, query)
|
||
|
|
|
||
|
|
chatReq := map[string]interface{}{
|
||
|
|
"model": model,
|
||
|
|
"messages": []interface{}{
|
||
|
|
map[string]interface{}{
|
||
|
|
"role": "user",
|
||
|
|
"content": context,
|
||
|
|
},
|
||
|
|
},
|
||
|
|
}
|
||
|
|
|
||
|
|
chatBody, _ := json.Marshal(chatReq)
|
||
|
|
chatHTTPReq, _ := http.NewRequest("POST", "/v1/chat/completions", io.NopCloser(bytes.NewReader(chatBody)))
|
||
|
|
chatHTTPReq.Header.Set("Content-Type", "application/json")
|
||
|
|
|
||
|
|
chatResp := &responseCapture{}
|
||
|
|
handler.ServeHTTP(chatResp, chatHTTPReq)
|
||
|
|
|
||
|
|
var chatResult map[string]interface{}
|
||
|
|
if err := json.Unmarshal(chatResp.body.Bytes(), &chatResult); err != nil {
|
||
|
|
return nil, fmt.Errorf("failed to parse chat response: %v", err)
|
||
|
|
}
|
||
|
|
|
||
|
|
return map[string]interface{}{
|
||
|
|
"reranked_documents": topDocs,
|
||
|
|
"chat_response": chatResult,
|
||
|
|
}, nil
|
||
|
|
}
|
||
|
|
|
||
|
|
// batchEmbeddingsWorkflow: Generate embeddings for multiple texts
|
||
|
|
func (h *Handler) batchEmbeddingsWorkflow(r *http.Request, handler *Handler, input map[string]interface{}) (interface{}, error) {
|
||
|
|
texts, ok := input["texts"].([]interface{})
|
||
|
|
if !ok || len(texts) == 0 {
|
||
|
|
return nil, fmt.Errorf("missing required parameter: texts (array)")
|
||
|
|
}
|
||
|
|
|
||
|
|
model, ok := input["model"].(string)
|
||
|
|
if !ok {
|
||
|
|
model = "nomic-ai/nomic-embed-text-v2-moe"
|
||
|
|
}
|
||
|
|
|
||
|
|
// Convert interface{} to []string
|
||
|
|
textStrings := make([]string, 0)
|
||
|
|
for _, t := range texts {
|
||
|
|
if str, ok := t.(string); ok {
|
||
|
|
textStrings = append(textStrings, str)
|
||
|
|
}
|
||
|
|
}
|
||
|
|
|
||
|
|
if len(textStrings) == 0 {
|
||
|
|
return nil, fmt.Errorf("no valid text strings in texts array")
|
||
|
|
}
|
||
|
|
|
||
|
|
embedReq := map[string]interface{}{
|
||
|
|
"model": model,
|
||
|
|
"input": textStrings,
|
||
|
|
}
|
||
|
|
|
||
|
|
embedBody, _ := json.Marshal(embedReq)
|
||
|
|
embedHTTPReq, _ := http.NewRequest("POST", "/v1/embeddings", io.NopCloser(bytes.NewReader(embedBody)))
|
||
|
|
embedHTTPReq.Header.Set("Content-Type", "application/json")
|
||
|
|
|
||
|
|
embedResp := &responseCapture{}
|
||
|
|
handler.ServeHTTP(embedResp, embedHTTPReq)
|
||
|
|
|
||
|
|
var embedResult map[string]interface{}
|
||
|
|
if err := json.Unmarshal(embedResp.body.Bytes(), &embedResult); err != nil {
|
||
|
|
return nil, fmt.Errorf("failed to parse embedding response: %v", err)
|
||
|
|
}
|
||
|
|
|
||
|
|
return embedResult, nil
|
||
|
|
}
|
||
|
|
|
||
|
|
// Utility functions
|
||
|
|
|
||
|
|
// responseCapture captures HTTP response for reuse within workflows
|
||
|
|
type responseCapture struct {
|
||
|
|
status int
|
||
|
|
header http.Header
|
||
|
|
body bytes.Buffer
|
||
|
|
}
|
||
|
|
|
||
|
|
func (w *responseCapture) Header() http.Header {
|
||
|
|
if w.header == nil {
|
||
|
|
w.header = make(http.Header)
|
||
|
|
}
|
||
|
|
return w.header
|
||
|
|
}
|
||
|
|
|
||
|
|
func (w *responseCapture) Write(b []byte) (int, error) {
|
||
|
|
if w.status == 0 {
|
||
|
|
w.status = http.StatusOK
|
||
|
|
}
|
||
|
|
return w.body.Write(b)
|
||
|
|
}
|
||
|
|
|
||
|
|
func (w *responseCapture) WriteHeader(statusCode int) {
|
||
|
|
if w.status == 0 {
|
||
|
|
w.status = statusCode
|
||
|
|
}
|
||
|
|
}
|
||
|
|
|
||
|
|
// generateWorkflowID generates a unique workflow execution ID
|
||
|
|
func generateWorkflowID() string {
|
||
|
|
return fmt.Sprintf("wf_%d", time.Now().UnixNano())
|
||
|
|
}
|
||
|
|
|
||
|
|
|