326 lines
11 KiB
Go
326 lines
11 KiB
Go
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()
|
|
}
|