Timothy Pomeroy преди 4 дни
родител
ревизия
f93efb56cb
променени са 2 файла, в които са добавени 177 реда и са изтрити 0 реда
  1. 15 0
      config.json
  2. 162 0
      prompt-rewriter.go

+ 15 - 0
config.json

@@ -0,0 +1,15 @@
+{
+  "plugins": [
+    {
+      "name": "prompt-rewriter",
+      "enabled": true,
+      "path": "prompt-rewriter.so",
+      "config": {
+        "small_model_endpoint": "https://bifrost.home.timandjenni.com/openai/v1/chat/completions",
+        "small_model_name": "qwen3-4b-instruct",
+        "enable_logging": true,
+        "timeout_seconds": 15
+      }
+    }
+  ]
+}

+ 162 - 0
prompt-rewriter.go

@@ -0,0 +1,162 @@
+package main
+
+import (
+	"bytes"
+	"encoding/json"
+	"fmt"
+	"io"
+	"net/http"
+	"strings"
+	"time"
+
+	"github.com/maximhq/bifrost/core/schemas"
+)
+
+// Config that can be set in Bifrost's config.json
+type PluginConfig struct {
+	SmallModelEndpoint string `json:"small_model_endpoint"` // e.g. "http://localhost:8081/v1/chat/completions"
+	SmallModelName     string `json:"small_model_name"`     // e.g. "qwen3-4b-instruct"
+	RewriteSystemPrompt string `json:"rewrite_system_prompt"`
+	EnableLogging      bool   `json:"enable_logging"`
+	TimeoutSeconds     int    `json:"timeout_seconds"`
+}
+
+var config PluginConfig
+
+func Init(cfg any) error {
+	// Default values
+	config = PluginConfig{
+		SmallModelEndpoint:  "https://bifrost.home.timandjenni.com/openai/v1/chat/completions", // change to your llama.cpp / ollama / lemonade endpoint
+		SmallModelName:      "qwen3-4b-instruct",
+		RewriteSystemPrompt: `You are an expert prompt engineer. Rewrite the user's raw request into a clear, well-structured, high-quality prompt that will produce better results from a larger coding/reasoning model. 
+Keep the original intent. Make it more precise, add useful structure if helpful, and remove ambiguity. 
+Return ONLY the improved prompt — no explanations.`,
+		EnableLogging:  true,
+		TimeoutSeconds: 15,
+	}
+
+	// Override with config from Bifrost if provided
+	if cfgMap, ok := cfg.(map[string]any); ok {
+		if v, ok := cfgMap["small_model_endpoint"].(string); ok {
+			config.SmallModelEndpoint = v
+		}
+		if v, ok := cfgMap["small_model_name"].(string); ok {
+			config.SmallModelName = v
+		}
+		if v, ok := cfgMap["rewrite_system_prompt"].(string); ok {
+			config.RewriteSystemPrompt = v
+		}
+		if v, ok := cfgMap["enable_logging"].(bool); ok {
+			config.EnableLogging = v
+		}
+	}
+	return nil
+}
+
+func GetName() string {
+	return "prompt-rewriter"
+}
+
+func PreLLMHook(ctx *schemas.BifrostContext, req *schemas.BifrostRequest) (*schemas.BifrostRequest, *schemas.LLMPluginShortCircuit, error) {
+	// Only act on chat requests that have messages
+	if req.ChatRequest == nil || len(req.ChatRequest.Input) == 0 {
+		return req, nil, nil
+	}
+
+	// Find the last user message
+	var lastUserMsg *schemas.ChatMessage
+	var lastUserIdx int
+	for i := len(req.ChatRequest.Input) - 1; i >= 0; i-- {
+		if req.ChatRequest.Input[i].Role == "user" {
+			lastUserMsg = &req.ChatRequest.Input[i]
+			lastUserIdx = i
+			break
+		}
+	}
+	if lastUserMsg == nil || lastUserMsg.Content == nil || lastUserMsg.Content.ContentStr == nil {
+		return req, nil, nil
+	}
+
+	originalPrompt := *lastUserMsg.Content.ContentStr
+	if strings.TrimSpace(originalPrompt) == "" {
+		return req, nil, nil
+	}
+
+	if config.EnableLogging {
+		ctx.Log(schemas.LogLevelInfo, fmt.Sprintf("[prompt-rewriter] Original prompt: %s", truncate(originalPrompt, 120)))
+	}
+
+	// Call the small model to rewrite the prompt
+	improved, err := rewritePrompt(originalPrompt)
+	if err != nil {
+		ctx.Log(schemas.LogLevelWarn, fmt.Sprintf("[prompt-rewriter] Rewrite failed, using original: %v", err))
+		return req, nil, nil // fail open — keep original prompt
+	}
+
+	if config.EnableLogging {
+		ctx.Log(schemas.LogLevelInfo, fmt.Sprintf("[prompt-rewriter] Improved prompt: %s", truncate(improved, 120)))
+	}
+
+	// Replace the user message content with the improved version
+	req.ChatRequest.Input[lastUserIdx].Content = &schemas.ChatMessageContent{
+		ContentStr: &improved,
+	}
+
+	return req, nil, nil
+}
+
+func rewritePrompt(original string) (string, error) {
+	payload := map[string]any{
+		"model": config.SmallModelName,
+		"messages": []map[string]string{
+			{"role": "system", "content": config.RewriteSystemPrompt},
+			{"role": "user", "content": original},
+		},
+		"temperature": 0.3,
+		"max_tokens":  1024,
+	}
+
+	body, _ := json.Marshal(payload)
+	client := &http.Client{Timeout: time.Duration(config.TimeoutSeconds) * time.Second}
+
+	resp, err := client.Post(config.SmallModelEndpoint, "application/json", bytes.NewReader(body))
+	if err != nil {
+		return "", err
+	}
+	defer resp.Body.Close()
+
+	if resp.StatusCode != 200 {
+		b, _ := io.ReadAll(resp.Body)
+		return "", fmt.Errorf("small model returned %d: %s", resp.StatusCode, string(b))
+	}
+
+	var result struct {
+		Choices []struct {
+			Message struct {
+				Content string `json:"content"`
+			} `json:"message"`
+		} `json:"choices"`
+	}
+	if err := json.NewDecoder(resp.Body).Decode(&result); err != nil {
+		return "", err
+	}
+	if len(result.Choices) == 0 {
+		return "", fmt.Errorf("no choices returned")
+	}
+
+	return strings.TrimSpace(result.Choices[0].Message.Content), nil
+}
+
+func truncate(s string, max int) string {
+	if len(s) <= max {
+		return s
+	}
+	return s[:max] + "..."
+}
+
+// Required stubs
+func PreRequestHook(_ *schemas.BifrostContext, _ *schemas.BifrostRequest) error { return nil }
+func PostLLMHook(_ *schemas.BifrostContext, resp *schemas.BifrostResponse, err *schemas.BifrostError) (*schemas.BifrostResponse, *schemas.BifrostError, error) {
+	return resp, err, nil
+}
+func Cleanup() error { return nil }