Files

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 are dispatched automatically as a bounded config.write job.")
}
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()
}