فهرست منبع

feat(tool): add tools to stack gateways

Thomas Zhang 2 ماه پیش
والد
کامیت
dcd3f06759
3فایلهای تغییر یافته به همراه134 افزوده شده و 3 حذف شده
  1. 34 0
      internal/controller/stack_controller.go
  2. 36 2
      internal/reconciler/litellm.go
  3. 64 1
      pkg/templates/litellm/config.go

+ 34 - 0
internal/controller/stack_controller.go

@@ -330,6 +330,34 @@ func (r *StackReconciler) reconcileGateways(ctx context.Context, s *v1alpha1.Sta
 		}
 	}
 
+	etList := &v1alpha1.ExternalToolList{}
+	if err := r.List(ctx, etList, client.InNamespace(s.Namespace)); err != nil {
+		return fmt.Errorf("Failed to list external tools: %v", err)
+	}
+	for _, em := range etList.Items {
+		if em.Spec.StackRef.Name == s.Name {
+			dependencies = append(dependencies, corev1.ObjectReference{
+				Kind:      "ExternalTool",
+				Name:      em.Name,
+				Namespace: em.Namespace,
+			})
+		}
+	}
+
+	mtList := &v1alpha1.ManagedToolList{}
+	if err := r.List(ctx, mtList, client.InNamespace(s.Namespace)); err != nil {
+		return fmt.Errorf("Failed to list managed tools: %v", err)
+	}
+	for _, em := range mtList.Items {
+		if em.Spec.StackRef.Name == s.Name {
+			dependencies = append(dependencies, corev1.ObjectReference{
+				Kind:      "ManagedTool",
+				Name:      em.Name,
+				Namespace: em.Namespace,
+			})
+		}
+	}
+
 	if len(dependencies) == 0 {
 		return nil
 	}
@@ -353,6 +381,10 @@ func (r *StackReconciler) SetupWithManager(mgr ctrl.Manager) error {
 			stackName = typed.Spec.StackRef.Name
 		case *v1alpha1.ManagedModel:
 			stackName = typed.Spec.StackRef.Name
+		case *v1alpha1.ExternalTool:
+			stackName = typed.Spec.StackRef.Name
+		case *v1alpha1.ManagedTool:
+			stackName = typed.Spec.StackRef.Name
 		default:
 			return nil
 		}
@@ -377,6 +409,8 @@ func (r *StackReconciler) SetupWithManager(mgr ctrl.Manager) error {
 		Owns(&corev1.PersistentVolumeClaim{}).
 		Watches(&v1alpha1.ExternalModel{}, handler.EnqueueRequestsFromMapFunc(mapToStack)).
 		Watches(&v1alpha1.ManagedModel{}, handler.EnqueueRequestsFromMapFunc(mapToStack)).
+		Watches(&v1alpha1.ExternalTool{}, handler.EnqueueRequestsFromMapFunc(mapToStack)).
+		Watches(&v1alpha1.ManagedTool{}, handler.EnqueueRequestsFromMapFunc(mapToStack)).
 		Named("stack").
 		Complete(r)
 }

+ 36 - 2
internal/reconciler/litellm.go

@@ -60,6 +60,8 @@ func (r *LiteLLMReconciler) ReconcileComponent(ctx context.Context, tmpl *v1alph
 	deps := make([]client.Object, 0)
 	emList := make([]*v1alpha1.ExternalModel, 0)
 	mmList := make([]*v1alpha1.ManagedModel, 0)
+	etList := make([]*v1alpha1.ExternalTool, 0)
+	mtList := make([]*v1alpha1.ManagedTool, 0)
 	for _, dep := range r.component.Spec.Dependencies {
 		switch dep.Kind {
 		case "ExternalModel":
@@ -76,9 +78,23 @@ func (r *LiteLLMReconciler) ReconcileComponent(ctx context.Context, tmpl *v1alph
 			}
 			mmList = append(mmList, mm)
 			deps = append(deps, mm)
+		case "ExternalTool":
+			et := &v1alpha1.ExternalTool{}
+			if err := r.client.Get(ctx, client.ObjectKey{Name: dep.Name, Namespace: dep.Namespace}, et); err != nil {
+				return nil, fmt.Errorf("Failed to get ExternalTool %s/%s: %w", dep.Namespace, dep.Name, err)
+			}
+			etList = append(etList, et)
+			deps = append(deps, et)
+		case "ManagedTool":
+			mt := &v1alpha1.ManagedTool{}
+			if err := r.client.Get(ctx, client.ObjectKey{Name: dep.Name, Namespace: dep.Namespace}, mt); err != nil {
+				return nil, fmt.Errorf("Failed to get ManagedTool %s/%s: %w", dep.Namespace, dep.Name, err)
+			}
+			mtList = append(mtList, mt)
+			deps = append(deps, mt)
 		}
 	}
-	configHash, err := r.reconcileConfigMap(ctx, emList, mmList)
+	configHash, err := r.reconcileConfigMap(ctx, emList, mmList, etList, mtList)
 	if err != nil {
 		return nil, fmt.Errorf("Failed to reconcile LiteLLM config: %w", err)
 	}
@@ -91,6 +107,22 @@ func (r *LiteLLMReconciler) ReconcileComponent(ctx context.Context, tmpl *v1alph
 			authSpecs["ExternalModel"][em.Name] = em.Spec.Auth
 		}
 	}
+	for _, et := range etList {
+		if et.Spec.Auth != nil {
+			if _, ok := authSpecs["ExternalTool"]; !ok {
+				authSpecs["ExternalTool"] = make(map[string]*v1alpha1.AuthSpec)
+			}
+			authSpecs["ExternalTool"][et.Name] = et.Spec.Auth
+		}
+	}
+	for _, mt := range mtList {
+		if mt.Spec.Auth != nil {
+			if _, ok := authSpecs["ManagedTool"]; !ok {
+				authSpecs["ManagedTool"] = make(map[string]*v1alpha1.AuthSpec)
+			}
+			authSpecs["ManagedTool"][mt.Name] = mt.Spec.Auth
+		}
+	}
 	injectSecrets(tmpl, authSpecs)
 	if tmpl.Metadata.Annotations == nil {
 		tmpl.Metadata.Annotations = make(map[string]string)
@@ -104,7 +136,7 @@ func (r *LiteLLMReconciler) ReconcileComponent(ctx context.Context, tmpl *v1alph
 	return deps, nil
 }
 
-func (r *LiteLLMReconciler) reconcileConfigMap(ctx context.Context, emList []*v1alpha1.ExternalModel, mmList []*v1alpha1.ManagedModel) (string, error) {
+func (r *LiteLLMReconciler) reconcileConfigMap(ctx context.Context, emList []*v1alpha1.ExternalModel, mmList []*v1alpha1.ManagedModel, etList []*v1alpha1.ExternalTool, mtList []*v1alpha1.ManagedTool) (string, error) {
 	var o11yComp *v1alpha1.Component
 	if r.stack.Spec.Observability != nil && r.stack.Spec.Observability.Enabled {
 		o11yComp = &v1alpha1.Component{}
@@ -118,6 +150,8 @@ func (r *LiteLLMReconciler) reconcileConfigMap(ctx context.Context, emList []*v1
 		ObservabilityComponent: o11yComp,
 		ExternalModels:         emList,
 		ManagedModels:          mmList,
+		ExternalTools:          etList,
+		ManagedTools:           mtList,
 	}
 	config, err := configBuilder.BuildLiteLLMConfig()
 	if err != nil {

+ 64 - 1
pkg/templates/litellm/config.go

@@ -80,6 +80,8 @@ type LiteLLMConfigBuilder struct {
 	ObservabilityComponent *v1alpha1.Component
 	ExternalModels         []*v1alpha1.ExternalModel
 	ManagedModels          []*v1alpha1.ManagedModel
+	ExternalTools          []*v1alpha1.ExternalTool
+	ManagedTools           []*v1alpha1.ManagedTool
 }
 
 func (l *LiteLLMConfigBuilder) BuildLiteLLMConfig() (LiteLLMConfig, error) {
@@ -95,6 +97,18 @@ func (l *LiteLLMConfigBuilder) BuildLiteLLMConfig() (LiteLLMConfig, error) {
 		entry := buildManagedModelEntry(mm)
 		modelList = append(modelList, entry)
 	}
+	mcpServers := map[string]mcpServerEntry{}
+	for _, et := range l.ExternalTools {
+		entry := buildExternalToolEntry(et)
+		mcpServers[et.Spec.ToolName] = entry
+	}
+	for _, mt := range l.ManagedTools {
+		if mt.Status.Endpoint == "" {
+			continue
+		}
+		entry := buildManagedToolEntry(mt)
+		mcpServers[mt.Spec.ToolName] = entry
+	}
 	phoenixEnabled := phoenixEnabled(l.Stack)
 	successCB := []string{}
 	failureCB := []string{}
@@ -104,7 +118,7 @@ func (l *LiteLLMConfigBuilder) BuildLiteLLMConfig() (LiteLLMConfig, error) {
 	}
 	cfg := LiteLLMConfig{
 		ModelList:  modelList,
-		MCPServers: map[string]mcpServerEntry{},
+		MCPServers: mcpServers,
 		LiteLLMSettings: &liteLLMSettings{
 			SuccessCallback: successCB,
 			FailureCallback: failureCB,
@@ -218,3 +232,52 @@ func applyInferenceParams(params map[string]any, inf *v1alpha1.InferenceParamete
 		params[k] = v
 	}
 }
+
+// buildExternalToolEntry converts an ExternalTool into a LiteLLM mcp_server entry.
+func buildExternalToolEntry(et *v1alpha1.ExternalTool) mcpServerEntry {
+	return buildToolEntry("ExternalTool", et.Spec.ToolName, et.Spec.Endpoint, et.Spec.Transport, et.Spec.Auth)
+}
+
+// buildManagedToolEntry converts a ManagedTool into a LiteLLM mcp_server entry.
+func buildManagedToolEntry(mt *v1alpha1.ManagedTool) mcpServerEntry {
+	url := mt.Status.Endpoint
+	switch mt.Spec.Transport {
+	case "http":
+		url = fmt.Sprintf("%s/mcp", strings.TrimSuffix(mt.Status.Endpoint, "/"))
+	case "sse":
+		url = fmt.Sprintf("%s/sse", strings.TrimSuffix(mt.Status.Endpoint, "/"))
+	}
+	return buildToolEntry("ManagedTool", mt.Spec.ToolName, url, mt.Spec.Transport, mt.Spec.Auth)
+}
+
+func buildToolEntry(toolKind string, toolName string, endpoint string, transport string, auth *v1alpha1.AuthSpec) mcpServerEntry {
+	entry := mcpServerEntry{
+		URL:       endpoint,
+		Transport: transport,
+	}
+	if auth == nil {
+		return entry
+	}
+	if auth.BearerToken != nil {
+		entry.AuthType = "bearer_token"
+		entry.AuthValue = "os.environ/" + AuthEnvVarName(toolKind, toolName)
+	} else if auth.APIKey != nil {
+		headerName := auth.APIKey.HeaderName
+		if headerName == "" {
+			headerName = "X-Api-Key"
+		}
+		if entry.StaticHeaders == nil {
+			entry.StaticHeaders = make(map[string]string)
+		}
+		entry.StaticHeaders[headerName] = "os.environ/" + AuthEnvVarName(toolKind, toolName)
+	}
+
+	for _, h := range auth.Headers {
+		if entry.StaticHeaders == nil {
+			entry.StaticHeaders = make(map[string]string)
+		}
+		entry.StaticHeaders[h.Name] = "os.environ/" + AuthEnvVarName(toolKind, toolName+"_"+h.Name)
+	}
+
+	return entry
+}