功能修改
This commit is contained in:
@@ -0,0 +1,325 @@
|
||||
package service
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"fmt"
|
||||
"io"
|
||||
"net"
|
||||
"net/http"
|
||||
"net/url"
|
||||
"os"
|
||||
"regexp"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"browser.local/platform/domain"
|
||||
)
|
||||
|
||||
const maxAIProviderResponseBytes = 1024 * 1024
|
||||
|
||||
type AIProviderSecretResolver interface {
|
||||
Resolve(reference string) (string, error)
|
||||
}
|
||||
|
||||
type EnvironmentAIProviderSecretResolver struct{}
|
||||
|
||||
func (EnvironmentAIProviderSecretResolver) Resolve(reference string) (string, error) {
|
||||
reference = strings.TrimSpace(reference)
|
||||
if strings.HasPrefix(reference, "env://") {
|
||||
name := strings.TrimPrefix(reference, "env://")
|
||||
if !validEnvironmentName(name) {
|
||||
return "", errors.New("AI provider secret reference is invalid")
|
||||
}
|
||||
if value := os.Getenv(name); value != "" {
|
||||
return value, nil
|
||||
}
|
||||
return "", errors.New("AI provider secret is unavailable")
|
||||
}
|
||||
if strings.HasPrefix(reference, "secret://providers/") {
|
||||
name := strings.TrimPrefix(reference, "secret://providers/")
|
||||
name = strings.ToUpper(regexp.MustCompile(`[^A-Za-z0-9]+`).ReplaceAllString(name, "_"))
|
||||
name = strings.Trim(name, "_")
|
||||
if name == "" {
|
||||
return "", errors.New("AI provider secret reference is invalid")
|
||||
}
|
||||
if value := os.Getenv("PLATFORM_AI_PROVIDER_" + name + "_API_KEY"); value != "" {
|
||||
return value, nil
|
||||
}
|
||||
return "", errors.New("AI provider secret is unavailable")
|
||||
}
|
||||
return "", errors.New("AI provider secret backend is unsupported")
|
||||
}
|
||||
|
||||
type HTTPAIProviderClient struct {
|
||||
HTTPClient *http.Client
|
||||
SecretResolver AIProviderSecretResolver
|
||||
}
|
||||
|
||||
type openAIChatRequest struct {
|
||||
Model string `json:"model"`
|
||||
Messages []openAIChatMessage `json:"messages"`
|
||||
Temperature float64 `json:"temperature"`
|
||||
}
|
||||
|
||||
type openAIChatMessage struct {
|
||||
Role string `json:"role"`
|
||||
Content string `json:"content"`
|
||||
}
|
||||
|
||||
type openAIChatResponse struct {
|
||||
Choices []struct {
|
||||
Message openAIChatMessage `json:"message"`
|
||||
} `json:"choices"`
|
||||
Usage struct {
|
||||
PromptTokens int `json:"prompt_tokens"`
|
||||
CompletionTokens int `json:"completion_tokens"`
|
||||
} `json:"usage"`
|
||||
}
|
||||
|
||||
type claudeRequest struct {
|
||||
Model string `json:"model"`
|
||||
MaxTokens int `json:"max_tokens"`
|
||||
Messages []openAIChatMessage `json:"messages"`
|
||||
}
|
||||
|
||||
type claudeResponse struct {
|
||||
Content []struct {
|
||||
Type string `json:"type"`
|
||||
Text string `json:"text"`
|
||||
} `json:"content"`
|
||||
Usage struct {
|
||||
InputTokens int `json:"input_tokens"`
|
||||
OutputTokens int `json:"output_tokens"`
|
||||
} `json:"usage"`
|
||||
}
|
||||
|
||||
type geminiRequest struct {
|
||||
Contents []struct {
|
||||
Parts []struct {
|
||||
Text string `json:"text"`
|
||||
} `json:"parts"`
|
||||
} `json:"contents"`
|
||||
}
|
||||
|
||||
type geminiResponse struct {
|
||||
Candidates []struct {
|
||||
Content struct {
|
||||
Parts []struct {
|
||||
Text string `json:"text"`
|
||||
} `json:"parts"`
|
||||
} `json:"content"`
|
||||
} `json:"candidates"`
|
||||
UsageMetadata struct {
|
||||
PromptTokenCount int `json:"promptTokenCount"`
|
||||
CandidatesTokenCount int `json:"candidatesTokenCount"`
|
||||
} `json:"usageMetadata"`
|
||||
}
|
||||
|
||||
type structuredAIRecommendation struct {
|
||||
Recommendation string `json:"recommendation"`
|
||||
SuggestedConfig string `json:"suggestedConfig"`
|
||||
}
|
||||
|
||||
func (client HTTPAIProviderClient) Invoke(provider domain.AIProvider, request domain.AIInvocationRequest) (domain.AIProviderInvocationResult, error) {
|
||||
model := safeModel(request.Model, provider)
|
||||
endpoint, err := providerEndpoint(provider, model)
|
||||
if err != nil {
|
||||
return domain.AIProviderInvocationResult{}, err
|
||||
}
|
||||
secret := ""
|
||||
if provider.RelayMode != domain.AIRelayModeLocal {
|
||||
resolver := client.SecretResolver
|
||||
if resolver == nil {
|
||||
resolver = EnvironmentAIProviderSecretResolver{}
|
||||
}
|
||||
secret, err = resolver.Resolve(provider.APIKeyRef)
|
||||
if err != nil {
|
||||
return domain.AIProviderInvocationResult{}, errors.New("AI provider credential is unavailable")
|
||||
}
|
||||
}
|
||||
prompt := boundedProviderPrompt(request)
|
||||
body, err := providerRequestBody(provider.Kind, model, prompt)
|
||||
if err != nil {
|
||||
return domain.AIProviderInvocationResult{}, err
|
||||
}
|
||||
httpRequest, err := http.NewRequest(http.MethodPost, endpoint, bytes.NewReader(body))
|
||||
if err != nil {
|
||||
return domain.AIProviderInvocationResult{}, errors.New("AI provider request could not be created")
|
||||
}
|
||||
httpRequest.Header.Set("Content-Type", "application/json")
|
||||
setProviderAuthorization(httpRequest, provider.Kind, secret)
|
||||
httpClient := client.HTTPClient
|
||||
if httpClient == nil {
|
||||
timeout := time.Duration(provider.TimeoutMS) * time.Millisecond
|
||||
if timeout <= 0 || timeout > 2*time.Minute {
|
||||
timeout = 30 * time.Second
|
||||
}
|
||||
httpClient = &http.Client{Timeout: timeout}
|
||||
}
|
||||
response, err := httpClient.Do(httpRequest)
|
||||
if err != nil {
|
||||
return domain.AIProviderInvocationResult{}, errors.New("AI provider transport failed")
|
||||
}
|
||||
defer response.Body.Close()
|
||||
if response.StatusCode < 200 || response.StatusCode >= 300 {
|
||||
return domain.AIProviderInvocationResult{}, fmt.Errorf("AI provider returned status class %dxx", response.StatusCode/100)
|
||||
}
|
||||
payload, err := io.ReadAll(io.LimitReader(response.Body, maxAIProviderResponseBytes+1))
|
||||
if err != nil || len(payload) > maxAIProviderResponseBytes {
|
||||
return domain.AIProviderInvocationResult{}, errors.New("AI provider response was invalid or too large")
|
||||
}
|
||||
result, err := parseProviderResponse(provider.Kind, payload)
|
||||
if err != nil {
|
||||
return domain.AIProviderInvocationResult{}, errors.New("AI provider response was invalid")
|
||||
}
|
||||
result.Usage.ProviderID = provider.ID
|
||||
result.Usage.Model = model
|
||||
result.Usage.Mocked = false
|
||||
if request.Purpose == "config.suggest" || request.Purpose == "config.generate" {
|
||||
result = parseStructuredRecommendation(result)
|
||||
}
|
||||
return result, nil
|
||||
}
|
||||
|
||||
func (svc *CoreService) ConfigureAIProviderMode(mode string) error {
|
||||
switch strings.ToLower(strings.TrimSpace(mode)) {
|
||||
case "mock", "test", "local":
|
||||
svc.aiProviderClient = MockAIProviderClient{}
|
||||
case "", "live", "http":
|
||||
svc.aiProviderClient = HTTPAIProviderClient{SecretResolver: EnvironmentAIProviderSecretResolver{}}
|
||||
default:
|
||||
return validationError("AI provider mode must be live or mock")
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func providerEndpoint(provider domain.AIProvider, model string) (string, error) {
|
||||
base, err := url.Parse(strings.TrimRight(provider.BaseURL, "/"))
|
||||
if err != nil || base.Host == "" {
|
||||
return "", errors.New("AI provider endpoint is invalid")
|
||||
}
|
||||
host := base.Hostname()
|
||||
if base.Scheme != "https" && !(base.Scheme == "http" && provider.RelayMode == domain.AIRelayModeLocal && isLoopbackHost(host)) {
|
||||
return "", errors.New("AI provider endpoint requires HTTPS")
|
||||
}
|
||||
switch provider.Kind {
|
||||
case domain.AIProviderKindClaude:
|
||||
base.Path = strings.TrimRight(base.Path, "/") + "/messages"
|
||||
case domain.AIProviderKindGemini:
|
||||
base.Path = strings.TrimRight(base.Path, "/") + "/models/" + url.PathEscape(model) + ":generateContent"
|
||||
default:
|
||||
base.Path = strings.TrimRight(base.Path, "/") + "/chat/completions"
|
||||
}
|
||||
base.RawQuery = ""
|
||||
base.Fragment = ""
|
||||
return base.String(), nil
|
||||
}
|
||||
|
||||
func providerRequestBody(kind domain.AIProviderKind, model, prompt string) ([]byte, error) {
|
||||
switch kind {
|
||||
case domain.AIProviderKindClaude:
|
||||
return json.Marshal(claudeRequest{Model: model, MaxTokens: 2048, Messages: []openAIChatMessage{{Role: "user", Content: prompt}}})
|
||||
case domain.AIProviderKindGemini:
|
||||
body := geminiRequest{}
|
||||
body.Contents = append(body.Contents, struct {
|
||||
Parts []struct {
|
||||
Text string `json:"text"`
|
||||
} `json:"parts"`
|
||||
}{Parts: []struct {
|
||||
Text string `json:"text"`
|
||||
}{{Text: prompt}}})
|
||||
return json.Marshal(body)
|
||||
default:
|
||||
return json.Marshal(openAIChatRequest{Model: model, Messages: []openAIChatMessage{{Role: "user", Content: prompt}}, Temperature: 0.2})
|
||||
}
|
||||
}
|
||||
|
||||
func parseProviderResponse(kind domain.AIProviderKind, payload []byte) (domain.AIProviderInvocationResult, error) {
|
||||
switch kind {
|
||||
case domain.AIProviderKindClaude:
|
||||
var response claudeResponse
|
||||
if err := json.Unmarshal(payload, &response); err != nil || len(response.Content) == 0 || strings.TrimSpace(response.Content[0].Text) == "" {
|
||||
return domain.AIProviderInvocationResult{}, errors.New("invalid Claude response")
|
||||
}
|
||||
return domain.AIProviderInvocationResult{Recommendation: response.Content[0].Text, Usage: domain.AIInvocationUsage{InputTokens: response.Usage.InputTokens, OutputTokens: response.Usage.OutputTokens}}, nil
|
||||
case domain.AIProviderKindGemini:
|
||||
var response geminiResponse
|
||||
if err := json.Unmarshal(payload, &response); err != nil || len(response.Candidates) == 0 || len(response.Candidates[0].Content.Parts) == 0 || strings.TrimSpace(response.Candidates[0].Content.Parts[0].Text) == "" {
|
||||
return domain.AIProviderInvocationResult{}, errors.New("invalid Gemini response")
|
||||
}
|
||||
return domain.AIProviderInvocationResult{Recommendation: response.Candidates[0].Content.Parts[0].Text, Usage: domain.AIInvocationUsage{InputTokens: response.UsageMetadata.PromptTokenCount, OutputTokens: response.UsageMetadata.CandidatesTokenCount}}, nil
|
||||
default:
|
||||
var response openAIChatResponse
|
||||
if err := json.Unmarshal(payload, &response); err != nil || len(response.Choices) == 0 || strings.TrimSpace(response.Choices[0].Message.Content) == "" {
|
||||
return domain.AIProviderInvocationResult{}, errors.New("invalid chat completion response")
|
||||
}
|
||||
return domain.AIProviderInvocationResult{Recommendation: response.Choices[0].Message.Content, Usage: domain.AIInvocationUsage{InputTokens: response.Usage.PromptTokens, OutputTokens: response.Usage.CompletionTokens}}, nil
|
||||
}
|
||||
}
|
||||
|
||||
func parseStructuredRecommendation(result domain.AIProviderInvocationResult) domain.AIProviderInvocationResult {
|
||||
raw := strings.TrimSpace(result.Recommendation)
|
||||
raw = strings.TrimPrefix(raw, "```json")
|
||||
raw = strings.TrimPrefix(raw, "```")
|
||||
raw = strings.TrimSuffix(raw, "```")
|
||||
var structured structuredAIRecommendation
|
||||
if json.Unmarshal([]byte(strings.TrimSpace(raw)), &structured) == nil && strings.TrimSpace(structured.SuggestedConfig) != "" {
|
||||
result.Recommendation = strings.TrimSpace(structured.Recommendation)
|
||||
result.SuggestedConfig = structured.SuggestedConfig
|
||||
}
|
||||
return result
|
||||
}
|
||||
|
||||
func boundedProviderPrompt(request domain.AIInvocationRequest) string {
|
||||
var builder strings.Builder
|
||||
builder.WriteString("Purpose: ")
|
||||
builder.WriteString(request.Purpose)
|
||||
builder.WriteString("\nRequest: ")
|
||||
builder.WriteString(request.Prompt)
|
||||
if request.CurrentConfig != "" {
|
||||
builder.WriteString("\nCurrent configuration:\n")
|
||||
builder.WriteString(request.CurrentConfig)
|
||||
}
|
||||
if request.Purpose == "config.suggest" || request.Purpose == "config.generate" {
|
||||
builder.WriteString("\nReturn JSON with recommendation and suggestedConfig. Configuration changes require separate operator approval.")
|
||||
}
|
||||
return builder.String()
|
||||
}
|
||||
|
||||
func setProviderAuthorization(request *http.Request, kind domain.AIProviderKind, secret string) {
|
||||
if secret == "" {
|
||||
return
|
||||
}
|
||||
switch kind {
|
||||
case domain.AIProviderKindClaude:
|
||||
request.Header.Set("x-api-key", secret)
|
||||
request.Header.Set("anthropic-version", "2023-06-01")
|
||||
case domain.AIProviderKindGemini:
|
||||
request.Header.Set("x-goog-api-key", secret)
|
||||
default:
|
||||
request.Header.Set("Authorization", "Bearer "+secret)
|
||||
}
|
||||
}
|
||||
|
||||
func validEnvironmentName(value string) bool {
|
||||
if value == "" {
|
||||
return false
|
||||
}
|
||||
for _, char := range value {
|
||||
if (char >= 'A' && char <= 'Z') || (char >= '0' && char <= '9') || char == '_' {
|
||||
continue
|
||||
}
|
||||
return false
|
||||
}
|
||||
return true
|
||||
}
|
||||
|
||||
func isLoopbackHost(host string) bool {
|
||||
if strings.EqualFold(host, "localhost") {
|
||||
return true
|
||||
}
|
||||
ip := net.ParseIP(host)
|
||||
return ip != nil && ip.IsLoopback()
|
||||
}
|
||||
Reference in New Issue
Block a user