Merge remote-tracking branch 'upstream/dev' into dev

# Conflicts:
#	internal/cron/service.go
This commit is contained in:
Goon committed 2026-06-15 14:14:16 +07:00
commit 591d809779
39 files changed
+1705 -92

No files matched your search

+13
View File
@@ -397,6 +397,19 @@ func registerProvidersFromDB(registry *providers.Registry, provStore store.Provi
prov := providers.NewOpenAIProvider(p.Name, p.APIKey, base, store.BytePlusDefaultModel) prov := providers.NewOpenAIProvider(p.Name, p.APIKey, base, store.BytePlusDefaultModel)
prov.WithProviderType(p.ProviderType) prov.WithProviderType(p.ProviderType)
registry.RegisterForTenant(p.TenantID, prov) registry.RegisterForTenant(p.TenantID, prov)
case store.ProviderKimiCoding:
// Moonshot Kimi Coding requires a fixed User-Agent on every request.
// OpenAI-compatible wire shape otherwise.
base := p.APIBase
if base == "" {
base = store.KimiCodingDefaultAPIBase
}
prov := providers.NewOpenAIProvider(p.Name, p.APIKey, base, store.KimiCodingDefaultModel)
prov.WithProviderType(p.ProviderType)
prov.WithExtraHeaders(map[string]string{
"User-Agent": store.KimiCodingRequiredUserAgent,
})
registry.RegisterForTenant(p.TenantID, prov)
default: default:
prov := providers.NewOpenAIProvider(p.Name, p.APIKey, p.APIBase, "") prov := providers.NewOpenAIProvider(p.Name, p.APIKey, p.APIBase, "")
prov.WithProviderType(p.ProviderType) prov.WithProviderType(p.ProviderType)
+1
View File
@@ -131,6 +131,7 @@ func seedConfigForContext(ctx context.Context, sc store.SystemConfigStore, cfg *
set("tts.auto", cfg.Tts.Auto) set("tts.auto", cfg.Tts.Auto)
set("tts.mode", cfg.Tts.Mode) set("tts.mode", cfg.Tts.Mode)
setInt("tts.max_length", cfg.Tts.MaxLength) setInt("tts.max_length", cfg.Tts.MaxLength)
setInt("tts.timeout_ms", cfg.Tts.TimeoutMs)
// Cron // Cron
setInt("cron.max_retries", cfg.Cron.MaxRetries) setInt("cron.max_retries", cfg.Cron.MaxRetries)
+32 -10
View File
@@ -60,6 +60,20 @@ type response struct {
} }
func main() { func main() {
if len(os.Args) > 1 {
req := request{Action: os.Args[1]}
if len(os.Args) > 2 {
req.Package = os.Args[2]
}
resp := handleRequest(req)
out, _ := json.Marshal(resp)
fmt.Println(string(out))
if !resp.OK {
os.Exit(1)
}
return
}
slog.Info("pkg-helper: starting", "socket", socketPath, "protocol", "v2") slog.Info("pkg-helper: starting", "socket", socketPath, "protocol", "v2")
// Remove stale socket. // Remove stale socket.
@@ -208,14 +222,26 @@ func handleRequest(req request) response {
} }
} }
var runApkFunc = runApk
func runApk(args ...string) ([]byte, error) {
if os.Getuid() != 0 {
if sudo, err := exec.LookPath("sudo"); err == nil {
// Use -n to prevent blocking on interactive password prompts
sudoArgs := append([]string{"-n", "apk"}, args...)
return exec.Command(sudo, sudoArgs...).CombinedOutput()
}
}
return exec.Command("apk", args...).CombinedOutput()
}
func doInstall(pkg string) response { func doInstall(pkg string) response {
apkMutex.Lock() apkMutex.Lock()
defer apkMutex.Unlock() defer apkMutex.Unlock()
slog.Info("pkg-helper: installing", "package", pkg) slog.Info("pkg-helper: installing", "package", pkg)
cmd := exec.Command("apk", "add", "--no-cache", pkg) out, err := runApkFunc("add", "--no-cache", pkg)
out, err := cmd.CombinedOutput()
if err != nil { if err != nil {
msg, code := classifyApkOutput(string(out), err) msg, code := classifyApkOutput(string(out), err)
slog.Error("pkg-helper: install failed", "package", pkg, "error", msg, "code", code) slog.Error("pkg-helper: install failed", "package", pkg, "error", msg, "code", code)
@@ -233,8 +259,7 @@ func doUninstall(pkg string) response {
slog.Info("pkg-helper: uninstalling", "package", pkg) slog.Info("pkg-helper: uninstalling", "package", pkg)
cmd := exec.Command("apk", "del", pkg) out, err := runApkFunc("del", pkg)
out, err := cmd.CombinedOutput()
if err != nil { if err != nil {
msg, code := classifyApkOutput(string(out), err) msg, code := classifyApkOutput(string(out), err)
slog.Error("pkg-helper: uninstall failed", "package", pkg, "error", msg, "code", code) slog.Error("pkg-helper: uninstall failed", "package", pkg, "error", msg, "code", code)
@@ -255,8 +280,7 @@ func doUpgrade(pkg string) response {
slog.Info("pkg-helper: upgrading", "package", pkg) slog.Info("pkg-helper: upgrading", "package", pkg)
cmd := exec.Command("apk", "add", "-u", pkg) out, err := runApkFunc("add", "-u", pkg)
out, err := cmd.CombinedOutput()
if err != nil { if err != nil {
msg, code := classifyApkOutput(string(out), err) msg, code := classifyApkOutput(string(out), err)
slog.Error("pkg-helper: upgrade failed", "package", pkg, "error", msg, "code", code) slog.Error("pkg-helper: upgrade failed", "package", pkg, "error", msg, "code", code)
@@ -274,8 +298,7 @@ func doUpdateIndex() response {
slog.Info("pkg-helper: updating index") slog.Info("pkg-helper: updating index")
cmd := exec.Command("apk", "update") out, err := runApkFunc("update")
out, err := cmd.CombinedOutput()
if err != nil { if err != nil {
msg, code := classifyApkOutput(string(out), err) msg, code := classifyApkOutput(string(out), err)
slog.Warn("pkg-helper: update-index failed", "error", msg, "code", code) slog.Warn("pkg-helper: update-index failed", "error", msg, "code", code)
@@ -292,8 +315,7 @@ func doListOutdated() response {
apkMutex.Lock() apkMutex.Lock()
defer apkMutex.Unlock() defer apkMutex.Unlock()
cmd := exec.Command("apk", "version", "-l", "<") out, err := runApkFunc("version", "-l", "<")
out, err := cmd.CombinedOutput()
if err != nil { if err != nil {
msg, code := classifyApkOutput(string(out), err) msg, code := classifyApkOutput(string(out), err)
return response{Error: msg, Code: code} return response{Error: msg, Code: code}
+44
View File
@@ -12,6 +12,13 @@ import (
// Note: Command execution tests are not included here since apk is not available // Note: Command execution tests are not included here since apk is not available
// in unit test environments. Integration tests would handle actual execution. // in unit test environments. Integration tests would handle actual execution.
func TestHandleRequest(t *testing.T) { func TestHandleRequest(t *testing.T) {
// Mock runApkFunc to avoid running actual commands or triggering sudo
origRunApkFunc := runApkFunc
runApkFunc = func(args ...string) ([]byte, error) {
return []byte("mock output"), nil
}
defer func() { runApkFunc = origRunApkFunc }()
tests := []struct { tests := []struct {
name string name string
req request req request
@@ -154,6 +161,13 @@ func TestValidPkgName(t *testing.T) {
// TestHandleRequest_AllActionsValidated tests both install and uninstall actions. // TestHandleRequest_AllActionsValidated tests both install and uninstall actions.
func TestHandleRequest_AllActionsValidated(t *testing.T) { func TestHandleRequest_AllActionsValidated(t *testing.T) {
// Mock runApkFunc to avoid running actual commands or triggering sudo
origRunApkFunc := runApkFunc
runApkFunc = func(args ...string) ([]byte, error) {
return []byte("mock output"), nil
}
defer func() { runApkFunc = origRunApkFunc }()
tests := []struct { tests := []struct {
action string action string
}{ }{
@@ -365,6 +379,12 @@ func TestValidPkgNameRegex_Compliance(t *testing.T) {
// TestHandleRequest_ErrorMessages tests that error messages are clear. // TestHandleRequest_ErrorMessages tests that error messages are clear.
func TestHandleRequest_ErrorMessages(t *testing.T) { func TestHandleRequest_ErrorMessages(t *testing.T) {
origRunApkFunc := runApkFunc
runApkFunc = func(args ...string) ([]byte, error) {
return []byte("mock output"), nil
}
defer func() { runApkFunc = origRunApkFunc }()
tests := []struct { tests := []struct {
name string name string
req request req request
@@ -404,6 +424,12 @@ func TestHandleRequest_ErrorMessages(t *testing.T) {
// Note: Actual apk command execution will fail in test environment (no apk available), // Note: Actual apk command execution will fail in test environment (no apk available),
// but validation should pass. // but validation should pass.
func TestHandleRequest_SuccessPath(t *testing.T) { func TestHandleRequest_SuccessPath(t *testing.T) {
origRunApkFunc := runApkFunc
runApkFunc = func(args ...string) ([]byte, error) {
return []byte("mock output"), nil
}
defer func() { runApkFunc = origRunApkFunc }()
tests := []struct { tests := []struct {
action string action string
pkg string pkg string
@@ -438,6 +464,12 @@ func TestHandleRequest_SuccessPath(t *testing.T) {
// TestHandleRequest_UpgradeValidation verifies that the upgrade action uses // TestHandleRequest_UpgradeValidation verifies that the upgrade action uses
// the stricter validApkName regex (lowercase only, no @, no /). // the stricter validApkName regex (lowercase only, no @, no /).
func TestHandleRequest_UpgradeValidation(t *testing.T) { func TestHandleRequest_UpgradeValidation(t *testing.T) {
origRunApkFunc := runApkFunc
runApkFunc = func(args ...string) ([]byte, error) {
return []byte("mock output"), nil
}
defer func() { runApkFunc = origRunApkFunc }()
// Valid names for upgrade (lowercase apk grammar) // Valid names for upgrade (lowercase apk grammar)
valid := []string{ valid := []string{
"curl", "curl",
@@ -462,6 +494,12 @@ func TestHandleRequest_UpgradeValidation(t *testing.T) {
// TestHandleRequest_UpgradeInjectionPatterns verifies 5 injection patterns are rejected. // TestHandleRequest_UpgradeInjectionPatterns verifies 5 injection patterns are rejected.
func TestHandleRequest_UpgradeInjectionPatterns(t *testing.T) { func TestHandleRequest_UpgradeInjectionPatterns(t *testing.T) {
origRunApkFunc := runApkFunc
runApkFunc = func(args ...string) ([]byte, error) {
return []byte("mock output"), nil
}
defer func() { runApkFunc = origRunApkFunc }()
injections := []string{ injections := []string{
"-malicious", // leading hyphen "-malicious", // leading hyphen
"pkg;evil", // semicolon "pkg;evil", // semicolon
@@ -486,6 +524,12 @@ func TestHandleRequest_UpgradeInjectionPatterns(t *testing.T) {
// by legacy validPkgName for install/uninstall) is REJECTED by upgrade action // by legacy validPkgName for install/uninstall) is REJECTED by upgrade action
// via the stricter validApkName. // via the stricter validApkName.
func TestHandleRequest_UpgradeRejectsLegacySymbols(t *testing.T) { func TestHandleRequest_UpgradeRejectsLegacySymbols(t *testing.T) {
origRunApkFunc := runApkFunc
runApkFunc = func(args ...string) ([]byte, error) {
return []byte("mock output"), nil
}
defer func() { runApkFunc = origRunApkFunc }()
legacySymbols := []string{ legacySymbols := []string{
"pkg@edge", // @ accepted by validPkgName, rejected by validApkName "pkg@edge", // @ accepted by validPkgName, rejected by validApkName
"@scope/pkg", // npm scoped — rejected by validApkName "@scope/pkg", // npm scoped — rejected by validApkName
+2 -2
View File
@@ -207,9 +207,9 @@ var coreToolSummaries = map[string]string{
"session_status": "Show session status (model, tokens, compaction count)", "session_status": "Show session status (model, tokens, compaction count)",
"sessions_history": "Fetch message history for a session", "sessions_history": "Fetch message history for a session",
"sessions_send": "Send a message into another session", "sessions_send": "Send a message into another session",
"read_image": "Analyze images — call with path from <media:image> tags", "read_image": "Analyze images — call with path from <media:image> tags, or a direct HTTP/HTTPS URL via the 'url' parameter",
"read_audio": "Analyze audio — call with media_id from <media:audio> tags", "read_audio": "Analyze audio — call with media_id from <media:audio> tags",
"read_video": "Analyze video — call with media_id from <media:video> tags", "read_video": "Analyze video — call with media_id from <media:video> tags, or a direct HTTP/HTTPS URL via the 'url' parameter",
"create_video": "Generate videos from text descriptions using AI", "create_video": "Generate videos from text descriptions using AI",
"read_document": "Analyze documents (PDF, DOCX) from <media:document> tags. If fails, use a skill instead. Path is directly accessible", "read_document": "Analyze documents (PDF, DOCX) from <media:document> tags. If fails, use a skill instead. Path is directly accessible",
"create_image": "Generate images from text descriptions using AI", "create_image": "Generate images from text descriptions using AI",
+1
View File
@@ -87,6 +87,7 @@ func (c *Config) ApplySystemConfigs(configs map[string]string) {
str("tts.auto", &c.Tts.Auto) str("tts.auto", &c.Tts.Auto)
str("tts.mode", &c.Tts.Mode) str("tts.mode", &c.Tts.Mode)
integer("tts.max_length", &c.Tts.MaxLength) integer("tts.max_length", &c.Tts.MaxLength)
integer("tts.timeout_ms", &c.Tts.TimeoutMs)
// Cron // Cron
integer("cron.max_retries", &c.Cron.MaxRetries) integer("cron.max_retries", &c.Cron.MaxRetries)
+3 -3
View File
@@ -222,10 +222,10 @@ func (cs *Service) GetJob(jobID string) (*Job, bool) {
cs.mu.Lock() cs.mu.Lock()
defer cs.mu.Unlock() defer cs.mu.Unlock()
for i, job := range cs.store.Jobs { for _, job := range cs.store.Jobs {
if job.ID == jobID { if job.ID == jobID {
jobCopy := cs.store.Jobs[i] result := job
return &jobCopy, true return &result, true
} }
} }
return nil, false return nil, false
+64 -6
View File
@@ -206,6 +206,32 @@ func TestService_EnableJob_NotFound(t *testing.T) {
} }
} }
func TestService_GetJob_ReturnsSnapshot(t *testing.T) {
dir := t.TempDir()
storePath := filepath.Join(dir, "cron.json")
cs := NewService(storePath, nil)
interval := int64(60000)
job, err := cs.AddJob("snapshot-job", Schedule{Kind: "every", EveryMS: &interval}, "hello", false, "", "", "agent-1")
if err != nil {
t.Fatalf("AddJob error: %v", err)
}
found, ok := cs.GetJob(job.ID)
if !ok {
t.Fatal("job should exist")
}
found.State.LastStatus = "mutated"
again, ok := cs.GetJob(job.ID)
if !ok {
t.Fatal("job should still exist")
}
if again.State.LastStatus == "mutated" {
t.Fatal("GetJob should return a snapshot, not internal service state")
}
}
// --- At-schedule sets DeleteAfterRun --- // --- At-schedule sets DeleteAfterRun ---
func TestService_AddJob_AtSchedule_DeleteAfterRun(t *testing.T) { func TestService_AddJob_AtSchedule_DeleteAfterRun(t *testing.T) {
@@ -235,18 +261,50 @@ func TestService_StartStop_JobExecution(t *testing.T) {
cs := NewService(storePath, handler) cs := NewService(storePath, handler)
interval := int64(50) if err := cs.Start(); err != nil {
_, err := cs.AddJob("fast", Schedule{Kind: "every", EveryMS: &interval}, "tick", false, "", "", "") t.Fatalf("Start error: %v", err)
}
defer cs.Stop()
interval := int64(time.Hour / time.Millisecond)
job, err := cs.AddJob("fast", Schedule{Kind: "every", EveryMS: &interval}, "tick", false, "", "", "")
if err != nil { if err != nil {
t.Fatalf("AddJob error: %v", err) t.Fatalf("AddJob error: %v", err)
} }
if err := cs.Start(); err != nil { cs.mu.Lock()
t.Fatalf("Start error: %v", err) foundJob := false
for i := range cs.store.Jobs {
if cs.store.Jobs[i].ID == job.ID {
due := nowMS()
cs.store.Jobs[i].State.NextRunAtMS = &due
foundJob = true
break
}
} }
if !foundJob {
cs.mu.Unlock()
t.Fatalf("job %s not found in store", job.ID)
}
if err := cs.saveUnsafe(); err != nil {
cs.mu.Unlock()
t.Fatalf("save due job: %v", err)
}
cs.mu.Unlock()
// fast tick = 20ms; wait enough for several ticks + at least 1 due fire deadline := time.Now().Add(500 * time.Millisecond)
time.Sleep(120 * time.Millisecond) ran := false
for time.Now().Before(deadline) {
found, ok := cs.GetJob(job.ID)
if ok && found.State.LastRunAtMS != nil {
ran = true
break
}
time.Sleep(5 * time.Millisecond)
}
if !ran {
t.Fatal("expected persisted job execution before deadline")
}
cs.Stop() cs.Stop()
count := execCount.Load() count := execCount.Load()
+89 -4
View File
@@ -4,8 +4,11 @@ import (
"context" "context"
"encoding/json" "encoding/json"
"errors" "errors"
"fmt"
"log/slog" "log/slog"
"net"
"regexp" "regexp"
"strconv"
"strings" "strings"
"github.com/google/uuid" "github.com/google/uuid"
@@ -13,6 +16,7 @@ import (
"github.com/nextlevelbuilder/goclaw/internal/gateway" "github.com/nextlevelbuilder/goclaw/internal/gateway"
"github.com/nextlevelbuilder/goclaw/internal/i18n" "github.com/nextlevelbuilder/goclaw/internal/i18n"
"github.com/nextlevelbuilder/goclaw/internal/permissions" "github.com/nextlevelbuilder/goclaw/internal/permissions"
"github.com/nextlevelbuilder/goclaw/internal/security"
"github.com/nextlevelbuilder/goclaw/internal/store" "github.com/nextlevelbuilder/goclaw/internal/store"
"github.com/nextlevelbuilder/goclaw/pkg/protocol" "github.com/nextlevelbuilder/goclaw/pkg/protocol"
) )
@@ -69,9 +73,13 @@ type bitrixPortalView struct {
// Validation regexes. Kept package-level so they compile once and tests can // Validation regexes. Kept package-level so they compile once and tests can
// reference them directly. // reference them directly.
var ( var (
// Bitrix24 cloud portal hosts. Matches *.bitrix24.{com,eu,ru,de,fr,jp,in,kz,ua,by} // Bitrix24 cloud portal hosts. Matches *.bitrix24.{com,eu,ru,de,fr,jp,in,kz,ua,by,vn,tr,es,com.br,com.ar}
// plus self-hosted *.bitrix.info. Subdomain regex matches DNS label rules. // plus self-hosted *.bitrix.info. Subdomain regex matches DNS label rules.
bitrixDomainRegex = regexp.MustCompile(`^[a-z0-9]([a-z0-9-]{0,61}[a-z0-9])?\.(bitrix24\.(com|eu|ru|de|fr|jp|in|kz|ua|by)|bitrix\.info)$`) bitrixCloudDomainRegex = regexp.MustCompile(`^[a-z0-9]([a-z0-9-]{0,61}[a-z0-9])?\.(bitrix24\.(com|eu|ru|de|fr|jp|in|kz|ua|by|vn|tr|es|com\.br|com\.ar)|bitrix\.info)$`)
// Valid hostname regex for self-hosted Bitrix24 instances (custom domains).
// Accepts any valid FQDN or hostname with optional port (e.g. bx.example.com, portal.internal:8443).
selfHostedDomainRegex = regexp.MustCompile(`^[a-z0-9]([a-z0-9-]{0,61}[a-z0-9])?(\.[a-z0-9]([a-z0-9-]{0,61}[a-z0-9])?)*(\:\d+)?$`)
// Portal name: lowercase slug used in install state token + channel config // Portal name: lowercase slug used in install state token + channel config
// reference. Underscore allowed for legacy CLI-created portals. // reference. Underscore allowed for legacy CLI-created portals.
@@ -142,10 +150,18 @@ func (m *BitrixPortalsMethods) handleCreate(ctx context.Context, client *gateway
client.SendResponse(protocol.NewErrorResponse(req.ID, protocol.ErrInvalidRequest, i18n.T(locale, i18n.MsgInvalidRequest, "name: lowercase letters, digits, hyphen, underscore (2-64 chars)"))) client.SendResponse(protocol.NewErrorResponse(req.ID, protocol.ErrInvalidRequest, i18n.T(locale, i18n.MsgInvalidRequest, "name: lowercase letters, digits, hyphen, underscore (2-64 chars)")))
return return
} }
if !bitrixDomainRegex.MatchString(domain) { if !bitrixCloudDomainRegex.MatchString(domain) && !selfHostedDomainRegex.MatchString(domain) {
client.SendResponse(protocol.NewErrorResponse(req.ID, protocol.ErrInvalidRequest, i18n.T(locale, i18n.MsgInvalidRequest, "domain: must be *.bitrix24.{com,eu,ru,…} or *.bitrix.info"))) client.SendResponse(protocol.NewErrorResponse(req.ID, protocol.ErrInvalidRequest, i18n.T(locale, i18n.MsgInvalidRequest, "domain: must be a valid hostname (e.g. *.bitrix24.com, *.bitrix.info, or your self-hosted domain)")))
return return
} }
// SSRF + port validation for self-hosted domains (cloud domains are
// Bitrix-operated and implicitly trusted).
if !bitrixCloudDomainRegex.MatchString(domain) {
if err := validateSelfHostedDomain(domain); err != nil {
client.SendResponse(protocol.NewErrorResponse(req.ID, protocol.ErrInvalidRequest, i18n.T(locale, i18n.MsgInvalidRequest, "domain: "+err.Error())))
return
}
}
if clientID == "" || clientSecret == "" { if clientID == "" || clientSecret == "" {
client.SendResponse(protocol.NewErrorResponse(req.ID, protocol.ErrInvalidRequest, i18n.T(locale, i18n.MsgRequired, "client_id and client_secret"))) client.SendResponse(protocol.NewErrorResponse(req.ID, protocol.ErrInvalidRequest, i18n.T(locale, i18n.MsgRequired, "client_id and client_secret")))
return return
@@ -352,6 +368,75 @@ func portalRowToView(row store.BitrixPortalData) bitrixPortalView {
return v return v
} }
// lookupHost is the DNS resolver used by validateSelfHostedDomain.
// Replaced in tests to avoid real network calls and to exercise multi-IP
// SSRF bypass scenarios.
var lookupHost = net.LookupHost
// validateSelfHostedDomain checks a self-hosted Bitrix24 domain for SSRF
// risks and invalid port ranges. Cloud domains (*.bitrix24.*, *.bitrix.info)
// are Bitrix-operated and implicitly trusted — this function is only called
// for custom/self-hosted domains.
//
// Policy:
// - Rejects literal private/loopback/metadata IPs (127.x, 10.x, 192.168.x,
// 169.254.x, ::1, fc00::, etc.)
// - Rejects hostnames where ANY resolved IP is blocked (not just the first)
// - Rejects .localhost and .local TLDs (commonly used for local dev)
// - Validates port range 1-65535 when a port is specified
func validateSelfHostedDomain(domain string) error {
// Extract host and optional port.
host := domain
portStr := ""
if idx := strings.LastIndex(domain, ":"); idx != -1 {
host = domain[:idx]
portStr = domain[idx+1:]
}
// Validate port range if present.
if portStr != "" {
port, err := strconv.Atoi(portStr)
if err != nil || port < 1 || port > 65535 {
return fmt.Errorf("port must be 1-65535")
}
}
// Reject .localhost and .local TLDs (commonly used for local development).
lowerHost := strings.ToLower(host)
if strings.HasSuffix(lowerHost, ".localhost") || strings.HasSuffix(lowerHost, ".local") || lowerHost == "localhost" {
return fmt.Errorf("private/internal hostnames (localhost, .local, .localhost) are not allowed")
}
// If the host is a literal IP, check it against blocked CIDRs.
if ip := net.ParseIP(host); ip != nil {
if security.IsBlocked(ip) {
return fmt.Errorf("IP %s is in a blocked range (loopback/private/metadata)", ip)
}
return nil
}
// For hostnames, resolve and check ALL returned IPs. DNS result ordering
// is not a security boundary — a resolver may return a public IP first
// followed by a private one (e.g. split-horizon, fallback addresses).
addrs, err := lookupHost(host)
if err != nil {
return fmt.Errorf("cannot resolve hostname %q", host)
}
if len(addrs) == 0 {
return fmt.Errorf("hostname %q resolved to no addresses", host)
}
for _, addr := range addrs {
ip := net.ParseIP(addr)
if ip == nil {
return fmt.Errorf("resolved address %q is not a valid IP", addr)
}
if security.IsBlocked(ip) {
return fmt.Errorf("hostname %q resolved to blocked IP %s (loopback/private/metadata)", host, ip)
}
}
return nil
}
// isDuplicateKeyErr probes a store error for a UNIQUE violation. Kept as a // isDuplicateKeyErr probes a store error for a UNIQUE violation. Kept as a
// string substring match because the store interface doesn't expose typed // string substring match because the store interface doesn't expose typed
// duplicate errors and we want consistent behaviour between pg + sqlite // duplicate errors and we want consistent behaviour between pg + sqlite
+180 -10
View File
@@ -4,6 +4,7 @@ import (
"context" "context"
"encoding/json" "encoding/json"
"errors" "errors"
"fmt"
"strings" "strings"
"sync" "sync"
"testing" "testing"
@@ -366,17 +367,62 @@ func TestBitrixPortals_Create_HappyPath_ReturnsInstallURL(t *testing.T) {
func TestBitrixPortals_Create_InvalidDomain(t *testing.T) { func TestBitrixPortals_Create_InvalidDomain(t *testing.T) {
tid := uuid.New() tid := uuid.New()
m := NewBitrixPortalsMethods(newStubBitrixPortalStore(), newStubChannelInstanceStore(), gatewayURLFn("https://gw.example.com")) m := NewBitrixPortalsMethods(newStubBitrixPortalStore(), newStubChannelInstanceStore(), gatewayURLFn("https://gw.example.com"))
client, ch := gateway.NewCapturingTestClient(permissions.RoleAdmin, tid, "u", 4)
cases := []struct {
name string
domain string
}{
{"leading hyphen", "-invalid.com"},
{"scheme included", "https://example.com"},
{"path included", "example.com/path"},
{"space in domain", "exam ple.com"},
{"invalid port zero", "example.com:0"},
{"invalid port overflow", "example.com:99999"},
{"localhost", "localhost"},
{"localhost subdomain", "bx.localhost"},
{"local TLD", "bx.local"},
}
for _, tc := range cases {
t.Run(tc.name, func(t *testing.T) {
client, ch := gateway.NewCapturingTestClient(permissions.RoleAdmin, tid, "u", 4)
m.handleCreate(store.WithTenantID(context.Background(), tid), client, buildBitrixReq(t, protocol.MethodBitrixPortalsCreate, map[string]string{
"name": "p",
"domain": tc.domain,
"client_id": "x",
"client_secret": "y",
}))
resp := readResponse(t, ch)
if resp.Error == nil || resp.Error.Code != protocol.ErrInvalidRequest {
t.Errorf("expected INVALID_REQUEST for %q, got %+v", tc.domain, resp.Error)
}
})
}
}
func TestBitrixPortals_Create_SelfHostedDomain(t *testing.T) {
tid := uuid.MustParse("22222222-2222-2222-2222-222222222222")
pStore := newStubBitrixPortalStore()
m := NewBitrixPortalsMethods(pStore, newStubChannelInstanceStore(), gatewayURLFn("https://goclaw.tamgiac.com"))
// Use a cloud domain (bitrixCloudDomainRegex) for the happy-path test
// since it bypasses SSRF DNS validation. Self-hosted SSRF validation
// is covered by TestValidateSelfHostedDomain_* tests.
client, ch := gateway.NewCapturingTestClient(permissions.RoleAdmin, tid, "admin", 4)
m.handleCreate(store.WithTenantID(context.Background(), tid), client, buildBitrixReq(t, protocol.MethodBitrixPortalsCreate, map[string]string{ m.handleCreate(store.WithTenantID(context.Background(), tid), client, buildBitrixReq(t, protocol.MethodBitrixPortalsCreate, map[string]string{
"name": "p", "name": "myportal",
"domain": "not-a-bitrix-domain.com", "domain": "myportal.bitrix24.com",
"client_id": "x", "client_id": "local.abc",
"client_secret": "y", "client_secret": "secret123",
})) }))
resp := readResponse(t, ch) resp := readResponse(t, ch)
if resp.Error == nil || resp.Error.Code != protocol.ErrInvalidRequest { if resp.Error != nil {
t.Errorf("expected INVALID_REQUEST, got %+v", resp.Error) t.Fatalf("create with self-hosted domain failed: %+v", resp.Error)
}
result := resp.Payload.(map[string]any)
if result["domain"] != "myportal.bitrix24.com" {
t.Errorf("domain = %q, want myportal.bitrix24.com", result["domain"])
} }
} }
@@ -540,22 +586,51 @@ func TestBitrixDomainRegex(t *testing.T) {
"my-corp.bitrix24.eu", "my-corp.bitrix24.eu",
"a.bitrix24.com", "a.bitrix24.com",
"company.bitrix.info", "company.bitrix.info",
"mycorp.bitrix24.vn",
"portal.bitrix24.tr",
"miempresa.bitrix24.es",
"empresa.bitrix24.com.br",
} }
bad := []string{ bad := []string{
"tamgiac.bitrix24", "tamgiac.bitrix24",
"tamgiac.example.com",
"tamgiac.bitrix24.xx", "tamgiac.bitrix24.xx",
"-bad.bitrix24.com", "-bad.bitrix24.com",
"UPPER.bitrix24.com", // we lowercase before match "UPPER.bitrix24.com", // we lowercase before match
"a.b.bitrix24.com", // multi-level subdomain not allowed "a.b.bitrix24.com", // multi-level subdomain not allowed
} }
for _, d := range good { for _, d := range good {
if !bitrixDomainRegex.MatchString(d) { if !bitrixCloudDomainRegex.MatchString(d) {
t.Errorf("should accept %q", d) t.Errorf("should accept %q", d)
} }
} }
for _, d := range bad { for _, d := range bad {
if bitrixDomainRegex.MatchString(d) { if bitrixCloudDomainRegex.MatchString(d) {
t.Errorf("should reject %q", d)
}
}
}
func TestSelfHostedDomainRegex(t *testing.T) {
good := []string{
"bx.example.com",
"portal.internal",
"bitrix.mycompany.co.uk",
"portal.example.com:8443",
"bx.corp",
}
bad := []string{
"-bad.example.com",
"UPPER.example.com", // we lowercase before match
"",
"not a domain",
}
for _, d := range good {
if !selfHostedDomainRegex.MatchString(d) {
t.Errorf("should accept %q", d)
}
}
for _, d := range bad {
if selfHostedDomainRegex.MatchString(d) {
t.Errorf("should reject %q", d) t.Errorf("should reject %q", d)
} }
} }
@@ -578,6 +653,101 @@ func TestPortalNameRegex(t *testing.T) {
} }
} }
// ---------------------------------------------------------------------------
// SSRF + port validation for self-hosted domains
// ---------------------------------------------------------------------------
func TestValidateSelfHostedDomain_SSRF(t *testing.T) {
// These should all be rejected as SSRF risks.
blocked := []string{
"127.0.0.1",
"127.0.0.1:8080",
"::1",
"10.0.0.1",
"10.0.0.1:443",
"192.168.1.1",
"172.16.0.1",
"169.254.169.254", // cloud metadata
"0.0.0.0",
"localhost",
"bx.localhost",
"bx.local",
"portal.internal.localhost",
}
for _, d := range blocked {
if err := validateSelfHostedDomain(d); err == nil {
t.Errorf("should reject %q (SSRF risk)", d)
}
}
}
func TestValidateSelfHostedDomain_PortRange(t *testing.T) {
badPorts := []string{
"example.com:0",
"example.com:99999",
"example.com:-1",
"example.com:abc",
}
for _, d := range badPorts {
if err := validateSelfHostedDomain(d); err == nil {
t.Errorf("should reject %q (invalid port)", d)
}
}
}
func TestValidateSelfHostedDomain_StubbedDNS(t *testing.T) {
// Save and restore the real resolver.
orig := lookupHost
defer func() { lookupHost = orig }()
tests := []struct {
name string
host string
addrs []string
dnsErr error
wantErr bool
}{
{"single public IP", "public.example.com", []string{"8.8.8.8"}, nil, false},
{"public IP with port", "public.example.com", []string{"93.184.216.34"}, nil, false},
{"single private IP", "internal.example.com", []string{"10.0.0.1"}, nil, true},
{"DNS resolution failure", "unresolvable.example.com", nil, fmt.Errorf("no such host"), true},
{"no addresses returned", "empty.example.com", []string{}, nil, true},
{"invalid IP string", "bad.example.com", []string{"not-an-ip"}, nil, true},
// Regression: multi-IP SSRF bypass — public IP first, private IP second.
// DNS ordering is not a security boundary; must reject if ANY IP is blocked.
{"multi-ip public-then-private", "evil.example.com", []string{"8.8.8.8", "10.0.0.1"}, nil, true},
{"multi-ip private-then-public", "evil2.example.com", []string{"192.168.1.1", "8.8.4.4"}, nil, true},
{"multi-ip all public", "safe.example.com", []string{"8.8.8.8", "8.8.4.4"}, nil, false},
{"multi-ip metadata IP", "meta.example.com", []string{"8.8.8.8", "169.254.169.254"}, nil, true},
{"multi-ip IPv6 loopback", "ipv6evil.example.com", []string{"2001:db8::1", "::1"}, nil, true},
}
for _, tc := range tests {
t.Run(tc.name, func(t *testing.T) {
lookupHost = func(host string) ([]string, error) {
if host == tc.host {
return tc.addrs, tc.dnsErr
}
return nil, fmt.Errorf("unexpected host: %s", host)
}
domain := tc.host
// Add port for the "public IP with port" case.
if tc.name == "public IP with port" {
domain = tc.host + ":443"
}
err := validateSelfHostedDomain(domain)
if tc.wantErr && err == nil {
t.Errorf("expected error for %q, got nil", tc.name)
}
if !tc.wantErr && err != nil {
t.Errorf("expected no error for %q, got: %v", tc.name, err)
}
})
}
}
// TestIsDuplicateKeyErr covers the two backend error string shapes we map // TestIsDuplicateKeyErr covers the two backend error string shapes we map
// to ALREADY_EXISTS. // to ALREADY_EXISTS.
func TestIsDuplicateKeyErr(t *testing.T) { func TestIsDuplicateKeyErr(t *testing.T) {
+26 -7
View File
@@ -15,8 +15,8 @@ import (
// ModelInfo is a normalized model entry returned by the list-models endpoint. // ModelInfo is a normalized model entry returned by the list-models endpoint.
type ModelInfo struct { type ModelInfo struct {
ID string `json:"id"` ID string `json:"id"`
Name string `json:"name,omitempty"` Name string `json:"name,omitempty"`
Reasoning *providers.ReasoningCapability `json:"reasoning,omitempty"` Reasoning *providers.ReasoningCapability `json:"reasoning,omitempty"`
} }
@@ -109,11 +109,8 @@ func (h *ProvidersHandler) handleListProviderModels(w http.ResponseWriter, r *ht
models = minimaxModels() models = minimaxModels()
default: default:
// All other types use OpenAI-compatible /models endpoint // All other types use OpenAI-compatible /models endpoint
apiBase := strings.TrimRight(h.resolveAPIBase(p), "/") apiBase := openAIModelsAPIBase(p.ProviderType, h.resolveAPIBase(p))
if apiBase == "" { models, err = fetchOpenAIModels(ctx, apiBase, p.APIKey, openAIModelsExtraHeaders(p.ProviderType))
apiBase = "https://api.openai.com/v1"
}
models, err = fetchOpenAIModels(ctx, apiBase, p.APIKey)
} }
if err != nil { if err != nil {
@@ -126,6 +123,28 @@ func (h *ProvidersHandler) handleListProviderModels(w http.ResponseWriter, r *ht
respond(withReasoningCapabilities(models)) respond(withReasoningCapabilities(models))
} }
func openAIModelsAPIBase(providerType, apiBase string) string {
base := strings.TrimRight(apiBase, "/")
if base != "" {
return base
}
switch providerType {
case store.ProviderKimiCoding:
return store.KimiCodingDefaultAPIBase
default:
return "https://api.openai.com/v1"
}
}
func openAIModelsExtraHeaders(providerType string) map[string]string {
if providerType != store.ProviderKimiCoding {
return nil
}
return map[string]string{
"User-Agent": store.KimiCodingRequiredUserAgent,
}
}
func reasoningDefaultsForModels( func reasoningDefaultsForModels(
settings []byte, settings []byte,
models []ModelInfo, models []ModelInfo,
+4 -1
View File
@@ -91,12 +91,15 @@ func fetchGeminiModels(ctx context.Context, apiKey string) ([]ModelInfo, error)
} }
// fetchOpenAIModels calls an OpenAI-compatible /models endpoint. // fetchOpenAIModels calls an OpenAI-compatible /models endpoint.
func fetchOpenAIModels(ctx context.Context, apiBase, apiKey string) ([]ModelInfo, error) { func fetchOpenAIModels(ctx context.Context, apiBase, apiKey string, extraHeaders map[string]string) ([]ModelInfo, error) {
req, err := http.NewRequestWithContext(ctx, "GET", apiBase+"/models", nil) req, err := http.NewRequestWithContext(ctx, "GET", apiBase+"/models", nil)
if err != nil { if err != nil {
return nil, err return nil, err
} }
req.Header.Set("Authorization", "Bearer "+apiKey) req.Header.Set("Authorization", "Bearer "+apiKey)
for k, v := range extraHeaders {
req.Header.Set(k, v)
}
resp, err := http.DefaultClient.Do(req) resp, err := http.DefaultClient.Do(req)
if err != nil { if err != nil {
+67
View File
@@ -227,6 +227,73 @@ func TestProvidersHandlerListProviderModelsOpenAICompatAnnotatesKnownModels(t *t
} }
} }
func TestProvidersHandlerListProviderModelsKimiCodingSendsRequiredUserAgent(t *testing.T) {
token := setupProvidersAdminToken(t)
var capturedAuth, capturedUserAgent string
upstream := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
capturedAuth = r.Header.Get("Authorization")
capturedUserAgent = r.Header.Get("User-Agent")
if r.URL.Path != "/models" {
http.NotFound(w, r)
return
}
_ = json.NewEncoder(w).Encode(map[string]any{
"data": []map[string]string{
{"id": store.KimiCodingDefaultModel},
},
})
}))
t.Cleanup(upstream.Close)
providerStore := newMockProviderStore()
provider := &store.LLMProviderData{
BaseModel: store.BaseModel{ID: uuid.New()},
Name: "kimi-coding",
ProviderType: store.ProviderKimiCoding,
APIBase: upstream.URL,
APIKey: "kimi-key",
Enabled: true,
}
if err := providerStore.CreateProvider(t.Context(), provider); err != nil {
t.Fatalf("CreateProvider() error = %v", err)
}
handler := NewProvidersHandler(providerStore, newMockSecretsStore(), nil, "")
mux := http.NewServeMux()
handler.RegisterRoutes(mux)
req := httptest.NewRequest(http.MethodGet, "/v1/providers/"+provider.ID.String()+"/models", nil)
req.Header.Set("Authorization", "Bearer "+token)
w := httptest.NewRecorder()
mux.ServeHTTP(w, req)
if w.Code != http.StatusOK {
t.Fatalf("status code = %d, want %d, body=%s", w.Code, http.StatusOK, w.Body.String())
}
if capturedAuth != "Bearer kimi-key" {
t.Fatalf("Authorization = %q, want Bearer kimi-key", capturedAuth)
}
if capturedUserAgent != store.KimiCodingRequiredUserAgent {
t.Fatalf("User-Agent = %q, want %q", capturedUserAgent, store.KimiCodingRequiredUserAgent)
}
var result ProviderModelsResponse
if err := json.NewDecoder(w.Body).Decode(&result); err != nil {
t.Fatalf("Decode() error = %v", err)
}
if len(result.Models) != 1 || result.Models[0].ID != store.KimiCodingDefaultModel {
t.Fatalf("models = %#v, want %q", result.Models, store.KimiCodingDefaultModel)
}
}
func TestOpenAIModelsAPIBaseDefaultsKimiCoding(t *testing.T) {
if got := openAIModelsAPIBase(store.ProviderKimiCoding, ""); got != store.KimiCodingDefaultAPIBase {
t.Fatalf("Kimi default api base = %q, want %q", got, store.KimiCodingDefaultAPIBase)
}
if got := openAIModelsAPIBase(store.ProviderOpenAICompat, ""); got != "https://api.openai.com/v1" {
t.Fatalf("OpenAI compat default api base = %q", got)
}
}
// TestProvidersHandlerListProviderModelsOllamaRichMetadata verifies that the // TestProvidersHandlerListProviderModelsOllamaRichMetadata verifies that the
// handler fetches /api/tags from Ollama and maps rich details (family, // handler fetches /api/tags from Ollama and maps rich details (family,
// parameter_size, quantization_level) into the display name. // parameter_size, quantization_level) into the display name.
+12
View File
@@ -314,6 +314,18 @@ func (h *ProvidersHandler) registerInMemory(p *store.LLMProviderData) providerRu
base = store.NovitaDefaultAPIBase base = store.NovitaDefaultAPIBase
} }
h.providerReg.RegisterForTenant(p.TenantID, providers.NewOpenAIProvider(p.Name, p.APIKey, base, store.NovitaDefaultModel)) h.providerReg.RegisterForTenant(p.TenantID, providers.NewOpenAIProvider(p.Name, p.APIKey, base, store.NovitaDefaultModel))
case store.ProviderKimiCoding:
// Moonshot Kimi Coding requires a fixed User-Agent on every request.
base := apiBase
if base == "" {
base = store.KimiCodingDefaultAPIBase
}
prov := providers.NewOpenAIProvider(p.Name, p.APIKey, base, store.KimiCodingDefaultModel)
prov.WithProviderType(p.ProviderType)
prov.WithExtraHeaders(map[string]string{
"User-Agent": store.KimiCodingRequiredUserAgent,
})
h.providerReg.RegisterForTenant(p.TenantID, prov)
default: default:
prov := providers.NewOpenAIProvider(p.Name, p.APIKey, apiBase, "") prov := providers.NewOpenAIProvider(p.Name, p.APIKey, apiBase, "")
if p.ProviderType == store.ProviderMiniMax { if p.ProviderType == store.ProviderMiniMax {
+4
View File
@@ -61,6 +61,10 @@ func (a *OpenAIAdapter) ToRequest(req ChatRequest) ([]byte, http.Header, error)
if a.provider.siteTitle != "" { if a.provider.siteTitle != "" {
h.Set("X-Title", a.provider.siteTitle) h.Set("X-Title", a.provider.siteTitle)
} }
// Mirror doRequest: provider-static headers (e.g. kimi_coding User-Agent).
for k, v := range a.provider.extraHeaders {
h.Set(k, v)
}
return data, h, nil return data, h, nil
} }
+2 -1
View File
@@ -98,7 +98,7 @@ func (p *AnthropicProvider) buildRequestBody(model string, req ChatRequest, stre
systemBlocks = append(systemBlocks, splitSystemPromptForCache(msg.Content)...) systemBlocks = append(systemBlocks, splitSystemPromptForCache(msg.Content)...)
case "user": case "user":
if len(msg.Images) > 0 { if len(msg.Images) > 0 || len(msg.Videos) > 0 {
var blocks []map[string]any var blocks []map[string]any
for _, img := range msg.Images { for _, img := range msg.Images {
blocks = append(blocks, map[string]any{ blocks = append(blocks, map[string]any{
@@ -110,6 +110,7 @@ func (p *AnthropicProvider) buildRequestBody(model string, req ChatRequest, stre
}, },
}) })
} }
// Videos are not supported by Anthropic, they are omitted here.
if msg.Content != "" { if msg.Content != "" {
blocks = append(blocks, map[string]any{ blocks = append(blocks, map[string]any{
"type": "text", "type": "text",
+2 -1
View File
@@ -34,7 +34,7 @@ func (p *CodexProvider) buildRequestBody(req ChatRequest, stream bool) map[strin
} }
case "user": case "user":
if len(m.Images) > 0 { if len(m.Images) > 0 || len(m.Videos) > 0 {
var parts []map[string]any var parts []map[string]any
for _, img := range m.Images { for _, img := range m.Images {
parts = append(parts, map[string]any{ parts = append(parts, map[string]any{
@@ -42,6 +42,7 @@ func (p *CodexProvider) buildRequestBody(req ChatRequest, stream bool) map[strin
"image_url": fmt.Sprintf("data:%s;base64,%s", img.MimeType, img.Data), "image_url": fmt.Sprintf("data:%s;base64,%s", img.MimeType, img.Data),
}) })
} }
// Videos are not supported by Codex, they are omitted here.
if m.Content != "" { if m.Content != "" {
parts = append(parts, map[string]any{ parts = append(parts, map[string]any{
"type": "input_text", "type": "input_text",
+31
View File
@@ -17,6 +17,7 @@ type OpenAIProvider struct {
providerType string // DB provider_type (e.g. "gemini_native", "openai", "minimax_native") providerType string // DB provider_type (e.g. "gemini_native", "openai", "minimax_native")
siteURL string // optional site URL for provider identification (e.g. OpenRouter HTTP-Referer) siteURL string // optional site URL for provider identification (e.g. OpenRouter HTTP-Referer)
siteTitle string // optional site title for provider identification (e.g. OpenRouter X-Title) siteTitle string // optional site title for provider identification (e.g. OpenRouter X-Title)
extraHeaders map[string]string // static headers set on every outgoing request (e.g. fixed User-Agent for kimi_coding)
client *http.Client client *http.Client
retryConfig RetryConfig retryConfig RetryConfig
middlewares RequestMiddleware // composed middleware chain (nil = no-op) middlewares RequestMiddleware // composed middleware chain (nil = no-op)
@@ -63,6 +64,36 @@ func (p *OpenAIProvider) WithSiteInfo(url, title string) *OpenAIProvider {
return p return p
} }
// WithExtraHeaders sets static headers attached to every outgoing request.
// Used by providers that require a fixed identity header (e.g. kimi_coding's
// User-Agent: claude-code/0.1.0). Repeat calls merge — keys already present are
// overwritten. Passing an empty map is a no-op.
func (p *OpenAIProvider) WithExtraHeaders(h map[string]string) *OpenAIProvider {
if len(h) == 0 {
return p
}
if p.extraHeaders == nil {
p.extraHeaders = make(map[string]string, len(h))
}
for k, v := range h {
p.extraHeaders[k] = v
}
return p
}
// ExtraHeaders returns a copy of the static headers configured for this provider.
// Used by adapter_openai.go to mirror the runtime request headers.
func (p *OpenAIProvider) ExtraHeaders() map[string]string {
if len(p.extraHeaders) == 0 {
return nil
}
out := make(map[string]string, len(p.extraHeaders))
for k, v := range p.extraHeaders {
out[k] = v
}
return out
}
// WithRegistry sets the model registry for forward-compat resolution. // WithRegistry sets the model registry for forward-compat resolution.
func (p *OpenAIProvider) WithRegistry(r ModelRegistry) *OpenAIProvider { func (p *OpenAIProvider) WithRegistry(r ModelRegistry) *OpenAIProvider {
p.registry = r p.registry = r
@@ -0,0 +1,203 @@
package providers
// Coverage for OpenAIProvider.WithExtraHeaders — the mechanism Kimi Coding
// uses to send a fixed User-Agent on every request.
import (
"context"
"io"
"net/http"
"net/http/httptest"
"testing"
)
// TestOpenAIProvider_ExtraHeaders_AppliedOnHTTPRequest verifies that headers
// set via WithExtraHeaders reach the actual outgoing request — not just the
// adapter's header map.
func TestOpenAIProvider_ExtraHeaders_AppliedOnHTTPRequest(t *testing.T) {
var gotUserAgent, gotXTrace, gotAuth string
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
gotUserAgent = r.Header.Get("User-Agent")
gotXTrace = r.Header.Get("X-Trace-Id")
gotAuth = r.Header.Get("Authorization")
// Minimal non-stream response so doRequest returns cleanly.
w.Header().Set("Content-Type", "application/json")
_, _ = w.Write([]byte(`{"id":"x","choices":[{"index":0,"message":{"role":"assistant","content":""},"finish_reason":"stop"}]}`))
}))
defer srv.Close()
p := NewOpenAIProvider("kimi-coding-test", "sk-fake", srv.URL, "kimi-k2-turbo-preview").
WithExtraHeaders(map[string]string{
"User-Agent": "claude-code/0.1.0",
"X-Trace-Id": "abc",
})
body, err := p.doRequest(context.Background(), map[string]any{
"model": "kimi-k2-turbo-preview",
"messages": []map[string]string{{"role": "user", "content": "hi"}},
})
if err != nil {
t.Fatalf("doRequest: %v", err)
}
_, _ = io.Copy(io.Discard, body)
_ = body.Close()
if gotUserAgent != "claude-code/0.1.0" {
t.Errorf("User-Agent = %q, want %q", gotUserAgent, "claude-code/0.1.0")
}
if gotXTrace != "abc" {
t.Errorf("X-Trace-Id = %q, want %q", gotXTrace, "abc")
}
// Standard Bearer auth must still apply alongside extra headers.
if gotAuth != "Bearer sk-fake" {
t.Errorf("Authorization = %q, want %q", gotAuth, "Bearer sk-fake")
}
}
// TestOpenAIAdapter_ExtraHeaders_MirroredInToRequest verifies the adapter path
// emits the same extra headers as the direct doRequest path — important
// because some call sites use adapter.ToRequest to produce headers separately.
func TestOpenAIAdapter_ExtraHeaders_MirroredInToRequest(t *testing.T) {
p := NewOpenAIProvider("kimi-coding-test", "sk-fake", "https://api.kimi.com/coding/v1", "kimi-k2-turbo-preview").
WithExtraHeaders(map[string]string{
"User-Agent": "claude-code/0.1.0",
})
a := &OpenAIAdapter{provider: p}
_, headers, err := a.ToRequest(ChatRequest{
Messages: []Message{{Role: "user", Content: "hi"}},
})
if err != nil {
t.Fatalf("ToRequest: %v", err)
}
if got := headers.Get("User-Agent"); got != "claude-code/0.1.0" {
t.Errorf("adapter User-Agent = %q, want claude-code/0.1.0", got)
}
}
// TestOpenAIProvider_ExtraHeaders_NoOpWhenEmpty makes sure the
// WithExtraHeaders(nil) / WithExtraHeaders({}) calls leave the provider's
// state alone — protects against accidental nil-map allocations in callers
// that pass through optional config.
func TestOpenAIProvider_ExtraHeaders_NoOpWhenEmpty(t *testing.T) {
p := NewOpenAIProvider("x", "k", "https://example.com", "m").
WithExtraHeaders(nil).
WithExtraHeaders(map[string]string{})
if got := p.ExtraHeaders(); got != nil {
t.Errorf("ExtraHeaders after empty calls = %v, want nil", got)
}
}
// TestKimiCoding_TemperatureSkipped reproduces the upstream rejection
// `invalid temperature: only 1 is allowed for this model`. When the provider
// is kimi_coding, the request body must omit temperature entirely so the
// upstream applies its mandatory default.
func TestKimiCoding_TemperatureSkipped(t *testing.T) {
p := NewOpenAIProvider("kimi-coding", "sk-fake", "https://api.kimi.com/coding/v1", "kimi-k2-turbo-preview").
WithProviderType("kimi_coding")
body := p.buildRequestBody("kimi-k2-turbo-preview", ChatRequest{
Messages: []Message{{Role: "user", Content: "hi"}},
Options: map[string]any{OptTemperature: 0.7},
}, true)
if _, present := body["temperature"]; present {
t.Errorf("temperature must not be sent to kimi_coding; got body[temperature]=%v", body["temperature"])
}
}
// TestKimiCoding_TemperatureSentForOtherProviders is the negative control —
// without provider_type=kimi_coding, a temperature option still flows through.
func TestKimiCoding_TemperatureSentForOtherProviders(t *testing.T) {
p := NewOpenAIProvider("openai", "sk-fake", "https://api.openai.com/v1", "gpt-4o-mini")
body := p.buildRequestBody("gpt-4o-mini", ChatRequest{
Messages: []Message{{Role: "user", Content: "hi"}},
Options: map[string]any{OptTemperature: 0.7},
}, true)
got, ok := body["temperature"]
if !ok {
t.Fatal("temperature must be sent for non-kimi providers")
}
if got != 0.7 {
t.Errorf("temperature = %v, want 0.7", got)
}
}
// TestKimiCoding_ReasoningContentAlwaysPresentOnAssistant reproduces upstream
// "thinking is enabled but reasoning_content is missing in assistant tool call
// message" — when an assistant message has tool_calls but no captured Thinking,
// kimi_coding must still carry reasoning_content (empty string is fine).
func TestKimiCoding_ReasoningContentAlwaysPresentOnAssistant(t *testing.T) {
p := NewOpenAIProvider("kimi-coding", "sk", "https://api.kimi.com/coding/v1", "kimi-k2-turbo-preview").
WithProviderType("kimi_coding")
body := p.buildRequestBody("kimi-k2-turbo-preview", ChatRequest{
Messages: []Message{
{Role: "user", Content: "list pods"},
{Role: "assistant", ToolCalls: []ToolCall{{ID: "call_1", Name: "exec", Arguments: map[string]any{"cmd": "kubectl get pods"}}}},
{Role: "tool", Content: "...", ToolCallID: "call_1"},
},
}, true)
msgs, ok := body["messages"].([]map[string]any)
if !ok {
t.Fatalf("messages not []map[string]any: %T", body["messages"])
}
if len(msgs) != 3 {
t.Fatalf("expected 3 messages, got %d", len(msgs))
}
// Assistant tool-call message must carry reasoning_content key.
assistant := msgs[1]
rc, present := assistant["reasoning_content"]
if !present {
t.Fatalf("kimi_coding assistant tool-call message must include reasoning_content key; got %v", assistant)
}
if rc != "" {
t.Errorf("reasoning_content = %q, want empty string when Thinking unset", rc)
}
}
// TestKimiCoding_ReasoningContentPreservedWhenSet ensures the empty-string
// fallback doesn't clobber real captured thinking content.
// (Use a non-trailing assistant message — buildRequestBody strips trailing
// assistant prefills as a safety net for proxy providers.)
func TestKimiCoding_ReasoningContentPreservedWhenSet(t *testing.T) {
p := NewOpenAIProvider("kimi-coding", "sk", "https://api.kimi.com/coding/v1", "kimi-k2-turbo-preview").
WithProviderType("kimi_coding")
body := p.buildRequestBody("kimi-k2-turbo-preview", ChatRequest{
Messages: []Message{
{Role: "user", Content: "hi"},
{Role: "assistant", Content: "hello", Thinking: "the user said hi"},
{Role: "user", Content: "more"},
},
}, true)
msgs := body["messages"].([]map[string]any)
if got := msgs[1]["reasoning_content"]; got != "the user said hi" {
t.Errorf("reasoning_content = %q, want %q", got, "the user said hi")
}
}
// TestNonKimi_ReasoningContentNotAddedWhenEmpty is the negative control — for
// other providers in the allowlist (e.g. deepseek), an empty Thinking must NOT
// inject an empty reasoning_content key, preserving today's behavior.
func TestNonKimi_ReasoningContentNotAddedWhenEmpty(t *testing.T) {
p := NewOpenAIProvider("deepseek", "sk", "https://api.deepseek.com/v1", "deepseek-chat")
body := p.buildRequestBody("deepseek-chat", ChatRequest{
Messages: []Message{
{Role: "user", Content: "hi"},
{Role: "assistant", ToolCalls: []ToolCall{{ID: "call_1", Name: "exec", Arguments: map[string]any{}}}},
{Role: "tool", Content: "...", ToolCallID: "call_1"},
},
}, true)
msgs := body["messages"].([]map[string]any)
if _, present := msgs[1]["reasoning_content"]; present {
t.Error("non-kimi providers must not inject empty reasoning_content; key should be absent")
}
}
+5
View File
@@ -45,6 +45,11 @@ func (p *OpenAIProvider) doRequest(ctx context.Context, body any) (io.ReadCloser
if p.siteTitle != "" { if p.siteTitle != "" {
httpReq.Header.Set("X-Title", p.siteTitle) httpReq.Header.Set("X-Title", p.siteTitle)
} }
// Static per-provider headers (e.g. fixed User-Agent for kimi_coding).
// Applied after the standard headers so providers can override them if needed.
for k, v := range p.extraHeaders {
httpReq.Header.Set(k, v)
}
resp, err := p.client.Do(httpReq) resp, err := p.client.Do(httpReq)
if err != nil { if err != nil {
+39 -5
View File
@@ -58,15 +58,27 @@ func (p *OpenAIProvider) buildRequestBody(model string, req ChatRequest, stream
// Echo reasoning_content only for APIs/models that accept it on assistant history. // Echo reasoning_content only for APIs/models that accept it on assistant history.
// Together Qwen and many OpenAI-compat gateways reject unknown message fields → HTTP 400. // Together Qwen and many OpenAI-compat gateways reject unknown message fields → HTTP 400.
if m.Thinking != "" && m.Role == "assistant" && openAIWireAssistantReasoningContent(model) { //
msg["reasoning_content"] = m.Thinking // Kimi Coding is stricter: when its server-side thinking is on (always-on for
// kimi-k2-turbo-preview), assistant tool-call messages MUST carry
// reasoning_content even if empty — otherwise upstream returns 400 "thinking
// is enabled but reasoning_content is missing in assistant tool call message".
if m.Role == "assistant" && openAIWireAssistantReasoningContent(model) {
switch {
case m.Thinking != "":
msg["reasoning_content"] = m.Thinking
case p.providerType == "kimi_coding":
// Send empty string rather than omit the field — satisfies Kimi's
// "must be present" check without inventing reasoning content.
msg["reasoning_content"] = ""
}
} }
// Include content; omit empty content for assistant messages with tool_calls // Include content; omit empty content for assistant messages with tool_calls
// (Gemini rejects empty content → "must include at least one parts field"). // (Gemini rejects empty content → "must include at least one parts field").
if m.Role == "user" && len(m.Images) > 0 { if m.Role == "user" && (len(m.Images) > 0 || len(m.Videos) > 0) {
var parts []map[string]any var parts []map[string]any
// Text before images — Together / Qwen vision examples use this order; OpenAI accepts both. // Text before images/videos — Together / Qwen vision examples use this order; OpenAI accepts both.
if m.Content != "" { if m.Content != "" {
parts = append(parts, map[string]any{ parts = append(parts, map[string]any{
"type": "text", "type": "text",
@@ -74,10 +86,26 @@ func (p *OpenAIProvider) buildRequestBody(model string, req ChatRequest, stream
}) })
} }
for _, img := range m.Images { for _, img := range m.Images {
urlVal := img.URL
if urlVal == "" {
urlVal = fmt.Sprintf("data:%s;base64,%s", img.MimeType, img.Data)
}
parts = append(parts, map[string]any{ parts = append(parts, map[string]any{
"type": "image_url", "type": "image_url",
"image_url": map[string]any{ "image_url": map[string]any{
"url": fmt.Sprintf("data:%s;base64,%s", img.MimeType, img.Data), "url": urlVal,
},
})
}
for _, vid := range m.Videos {
urlVal := vid.URL
if urlVal == "" {
urlVal = fmt.Sprintf("data:%s;base64,%s", vid.MimeType, vid.Data)
}
parts = append(parts, map[string]any{
"type": "video_url",
"video_url": map[string]any{
"url": urlVal,
}, },
}) })
} }
@@ -188,6 +216,12 @@ func (p *OpenAIProvider) buildRequestBody(model string, req ChatRequest, stream
// Note: gpt-5.X flagship models (gpt-5.1, gpt-5.4, gpt-5.5) DO support temperature; // Note: gpt-5.X flagship models (gpt-5.1, gpt-5.4, gpt-5.5) DO support temperature;
// only the mini/nano reasoning variants reject it. // only the mini/nano reasoning variants reject it.
skipTemp := strings.HasPrefix(capabilityModel, "gpt-5-mini") || strings.HasPrefix(capabilityModel, "gpt-5-nano") || strings.HasPrefix(capabilityModel, "o1") || strings.HasPrefix(capabilityModel, "o3") || strings.HasPrefix(capabilityModel, "o4") skipTemp := strings.HasPrefix(capabilityModel, "gpt-5-mini") || strings.HasPrefix(capabilityModel, "gpt-5-nano") || strings.HasPrefix(capabilityModel, "o1") || strings.HasPrefix(capabilityModel, "o3") || strings.HasPrefix(capabilityModel, "o4")
// Kimi Coding rejects any temperature override — `invalid temperature: only
// 1 is allowed for this model`. Skip sending so the upstream applies its
// own default (1). Matches the model-locked behavior of o1/o3/o4.
if p.providerType == "kimi_coding" {
skipTemp = true
}
if !skipTemp { if !skipTemp {
body["temperature"] = v body["temperature"] = v
} }
+54
View File
@@ -240,6 +240,60 @@ func TestBuildRequestBody_MultimodalTextBeforeImages(t *testing.T) {
} }
} }
func TestBuildRequestBody_MultimodalWithImageURL(t *testing.T) {
p := NewOpenAIProvider("test", "key", "https://api.openai.com/v1", "gpt-4")
req := ChatRequest{
Messages: []Message{
{
Role: "user",
Content: "describe",
Images: []ImageContent{
{URL: "https://example.com/image.png"},
},
},
},
}
body := p.buildRequestBody("gpt-4o", req, false)
msgs := body["messages"].([]map[string]any)
parts, ok := msgs[0]["content"].([]map[string]any)
if !ok || len(parts) < 2 {
t.Fatalf("want multimodal parts, got %v", msgs[0]["content"])
}
imgPart := parts[1]["image_url"].(map[string]any)
if urlVal, _ := imgPart["url"].(string); urlVal != "https://example.com/image.png" {
t.Errorf("expected URL to be https://example.com/image.png, got %q", urlVal)
}
}
func TestBuildRequestBody_MultimodalWithVideoURL(t *testing.T) {
p := NewOpenAIProvider("test", "key", "https://api.openai.com/v1", "gpt-4")
req := ChatRequest{
Messages: []Message{
{
Role: "user",
Content: "describe video",
Videos: []VideoContent{
{MimeType: "video/mp4", URL: "https://example.com/video.mp4"},
},
},
},
}
body := p.buildRequestBody("gpt-4o", req, false)
msgs := body["messages"].([]map[string]any)
parts, ok := msgs[0]["content"].([]map[string]any)
if !ok || len(parts) < 2 {
t.Fatalf("want multimodal parts, got %v", msgs[0]["content"])
}
videoPart, ok := parts[1]["video_url"].(map[string]any)
if !ok {
t.Fatalf("expected video_url part, got %v", parts[1])
}
if urlVal, _ := videoPart["url"].(string); urlVal != "https://example.com/video.mp4" {
t.Errorf("expected URL to be https://example.com/video.mp4, got %q", urlVal)
}
}
func TestBuildRequestBody_TogetherDetectedByProviderType(t *testing.T) { func TestBuildRequestBody_TogetherDetectedByProviderType(t *testing.T) {
// Together behind reverse proxy — detected by providerType, not URL. // Together behind reverse proxy — detected by providerType, not URL.
p := NewOpenAIProvider("my-proxy", "key", "https://proxy.internal/v1", "") p := NewOpenAIProvider("my-proxy", "key", "https://proxy.internal/v1", "")
+11 -1
View File
@@ -113,13 +113,22 @@ type StreamChunk struct {
Images []ImageContent `json:"images,omitempty"` // image generation frames (Codex) Images []ImageContent `json:"images,omitempty"` // image generation frames (Codex)
} }
// ImageContent represents a base64-encoded image for vision-capable models. // ImageContent represents an image (either base64-encoded or a direct URL) for vision-capable models.
type ImageContent struct { type ImageContent struct {
MimeType string `json:"mime_type"` // e.g. "image/jpeg" MimeType string `json:"mime_type"` // e.g. "image/jpeg"
Data string `json:"data"` // base64-encoded image bytes Data string `json:"data"` // base64-encoded image bytes
URL string `json:"url,omitempty"` // URL of the image
Partial bool `json:"partial,omitempty"` // true for intermediate frames (Codex image_generation_call) Partial bool `json:"partial,omitempty"` // true for intermediate frames (Codex image_generation_call)
} }
// VideoContent represents a video (either base64-encoded or a direct URL) for video-capable models.
type VideoContent struct {
MimeType string `json:"mime_type"` // e.g. "video/mp4"
Data string `json:"data"` // base64-encoded video bytes
URL string `json:"url,omitempty"` // URL of the video
Partial bool `json:"partial,omitempty"` // true for intermediate frames
}
// MediaRef is a lightweight reference to a persistently stored media file. // MediaRef is a lightweight reference to a persistently stored media file.
// Stored in session JSONB (~60 bytes each) instead of megabytes for base64. // Stored in session JSONB (~60 bytes each) instead of megabytes for base64.
// On reload, MediaRefs are resolved to file paths and loaded into Images (for images). // On reload, MediaRefs are resolved to file paths and loaded into Images (for images).
@@ -137,6 +146,7 @@ type Message struct {
Content string `json:"content"` Content string `json:"content"`
Thinking string `json:"thinking,omitempty"` // reasoning_content for thinking models (Kimi, DeepSeek, etc.) Thinking string `json:"thinking,omitempty"` // reasoning_content for thinking models (Kimi, DeepSeek, etc.)
Images []ImageContent `json:"-"` // vision: base64 images (runtime only, never persisted to DB) Images []ImageContent `json:"-"` // vision: base64 images (runtime only, never persisted to DB)
Videos []VideoContent `json:"-"` // vision: base64 videos (runtime only, never persisted to DB)
MediaRefs []MediaRef `json:"media_refs,omitempty"` // persistent media file references MediaRefs []MediaRef `json:"media_refs,omitempty"` // persistent media file references
ToolCalls []ToolCall `json:"tool_calls,omitempty"` ToolCalls []ToolCall `json:"tool_calls,omitempty"`
ToolCallID string `json:"tool_call_id,omitempty"` // for role="tool" responses ToolCallID string `json:"tool_call_id,omitempty"` // for role="tool" responses
+104
View File
@@ -6,6 +6,8 @@ import (
"encoding/json" "encoding/json"
"fmt" "fmt"
"net" "net"
"os"
"path/filepath"
"runtime" "runtime"
"strings" "strings"
"sync/atomic" "sync/atomic"
@@ -132,6 +134,108 @@ func TestApkHelperCall_DialFail(t *testing.T) {
} }
} }
func TestApkHelperCallFallback_ParsesStdoutJSONWithStderrLogs(t *testing.T) {
if runtime.GOOS == "windows" {
t.Skip("shell fixture requires Unix")
}
helper := writePkgHelperFixture(t, `
echo 'time=2026-06-15 level=INFO msg="pkg-helper: installing"' >&2
printf '%s\n' '{"ok":true,"data":"installed"}'
`)
ok, code, data, errMsg := apkHelperCallFallback(context.Background(), helper, "install", "curl")
if !ok {
t.Fatalf("ok = false, want true (code=%q err=%q)", code, errMsg)
}
if code != "" {
t.Errorf("code = %q, want empty", code)
}
if data != "installed" {
t.Errorf("data = %q, want %q", data, "installed")
}
if errMsg != "" {
t.Errorf("errMsg = %q, want empty", errMsg)
}
}
func TestApkHelperCallFallback_ParsesErrorJSONDespiteExitStatus(t *testing.T) {
if runtime.GOOS == "windows" {
t.Skip("shell fixture requires Unix")
}
helper := writePkgHelperFixture(t, `
echo 'time=2026-06-15 level=ERROR msg="pkg-helper: install failed"' >&2
printf '%s\n' '{"ok":false,"error":"package not found","code":"not_found"}'
exit 1
`)
ok, code, _, errMsg := apkHelperCallFallback(context.Background(), helper, "install", "missing")
if ok {
t.Fatal("ok = true, want false")
}
if code != "not_found" {
t.Errorf("code = %q, want %q", code, "not_found")
}
if errMsg != "package not found" {
t.Errorf("errMsg = %q, want %q", errMsg, "package not found")
}
}
func TestApkHelperCallFallback_InvalidStdoutIsHelperError(t *testing.T) {
if runtime.GOOS == "windows" {
t.Skip("shell fixture requires Unix")
}
helper := writePkgHelperFixture(t, `
printf '%s\n' 'not-json'
`)
ok, code, _, errMsg := apkHelperCallFallback(context.Background(), helper, "install", "curl")
if ok {
t.Fatal("ok = true, want false for invalid helper response")
}
if code != "helper_error" {
t.Errorf("code = %q, want %q", code, "helper_error")
}
if !strings.Contains(errMsg, "invalid response") || !strings.Contains(errMsg, "stdout: not-json") {
t.Errorf("errMsg = %q, want invalid response with stdout detail", errMsg)
}
}
func TestFirstExecutableFileFindsBundledFallbackCandidate(t *testing.T) {
if runtime.GOOS == "windows" {
t.Skip("executable-bit fixture requires Unix")
}
helper := writePkgHelperFixture(t, `exit 0`)
got, ok := firstExecutableFile([]string{
filepath.Join(t.TempDir(), "missing-helper"),
helper,
})
if !ok {
t.Fatal("firstExecutableFile did not find executable fallback candidate")
}
if got != helper {
t.Fatalf("firstExecutableFile = %q, want %q", got, helper)
}
}
func writePkgHelperFixture(t *testing.T, body string) string {
t.Helper()
path := filepath.Join(t.TempDir(), "pkg-helper")
script := "#!/bin/sh\n" + strings.TrimLeft(body, "\n")
if err := os.WriteFile(path, []byte(script), 0o755); err != nil {
t.Fatalf("write helper fixture: %v", err)
}
return path
}
// TestApkHelperCall_ValidResponse verifies a well-formed canned response is // TestApkHelperCall_ValidResponse verifies a well-formed canned response is
// parsed correctly into (ok, code, data, errMsg). // parsed correctly into (ok, code, data, errMsg).
func TestApkHelperCall_ValidResponse(t *testing.T) { func TestApkHelperCall_ValidResponse(t *testing.T) {
+98 -8
View File
@@ -2,6 +2,7 @@ package skills
import ( import (
"bufio" "bufio"
"bytes"
"context" "context"
"encoding/json" "encoding/json"
"fmt" "fmt"
@@ -39,6 +40,10 @@ const InstallTimeout = 5 * time.Minute
// pkgHelperSocket is the Unix socket path for the root-privileged pkg-helper. // pkgHelperSocket is the Unix socket path for the root-privileged pkg-helper.
const pkgHelperSocket = "/tmp/pkg.sock" const pkgHelperSocket = "/tmp/pkg.sock"
// pkgHelperBundledPath is where the Docker image copies the helper binary.
// /app is not part of PATH, so direct-exec fallback must check it explicitly.
const pkgHelperBundledPath = "/app/pkg-helper"
// apkHelperCallFunc is the package-level hook for apkHelperCall, allowing tests // apkHelperCallFunc is the package-level hook for apkHelperCall, allowing tests
// to inject a stub without starting a real Unix socket server. Production code // to inject a stub without starting a real Unix socket server. Production code
// always uses the default value (apkHelperCall). Tests replace it per-case and // always uses the default value (apkHelperCall). Tests replace it per-case and
@@ -306,6 +311,10 @@ func UninstallPackage(ctx context.Context, dep string) (bool, string) {
func apkHelperCall(ctx context.Context, action, pkg string) (ok bool, code, data, errMsg string) { func apkHelperCall(ctx context.Context, action, pkg string) (ok bool, code, data, errMsg string) {
conn, err := net.DialTimeout("unix", pkgHelperSocket, 5*time.Second) conn, err := net.DialTimeout("unix", pkgHelperSocket, 5*time.Second)
if err != nil { if err != nil {
if path, found := findPkgHelperBinary(); found {
return apkHelperCallFallback(ctx, path, action, pkg)
}
return false, "helper_unavailable", "", fmt.Sprintf("pkg-helper unavailable: %v", err) return false, "helper_unavailable", "", fmt.Sprintf("pkg-helper unavailable: %v", err)
} }
defer conn.Close() defer conn.Close()
@@ -334,22 +343,103 @@ func apkHelperCall(ctx context.Context, action, pkg string) (ok bool, code, data
return false, "helper_error", "", "pkg-helper: no response" return false, "helper_error", "", "pkg-helper: no response"
} }
var resp struct { resp, err := parsePkgHelperResponse(scanner.Bytes())
OK bool `json:"ok"` if err != nil {
Error string `json:"error"`
Code string `json:"code"`
Data string `json:"data"`
}
if err := json.Unmarshal(scanner.Bytes(), &resp); err != nil {
return false, "helper_error", "", fmt.Sprintf("pkg-helper: invalid response: %v", err) return false, "helper_error", "", fmt.Sprintf("pkg-helper: invalid response: %v", err)
} }
return resp.OK, resp.Code, resp.Data, resp.Error
}
type pkgHelperResponse struct {
OK bool `json:"ok"`
Error string `json:"error"`
Code string `json:"code"`
Data string `json:"data"`
}
func parsePkgHelperResponse(out []byte) (pkgHelperResponse, error) {
var resp pkgHelperResponse
if err := json.Unmarshal(out, &resp); err != nil {
return pkgHelperResponse{}, err
}
// Default missing code to system_error for v1-era helpers that omit the field. // Default missing code to system_error for v1-era helpers that omit the field.
if resp.Code == "" && !resp.OK { if resp.Code == "" && !resp.OK {
resp.Code = "system_error" resp.Code = "system_error"
} }
return resp, nil
}
return resp.OK, resp.Code, resp.Data, resp.Error func apkHelperCallFallback(ctx context.Context, helperPath, action, pkg string) (ok bool, code, data, errMsg string) {
cmd := exec.CommandContext(ctx, helperPath, action)
if pkg != "" {
cmd.Args = append(cmd.Args, pkg)
}
var stderr bytes.Buffer
cmd.Stderr = &stderr
out, execErr := cmd.Output()
resp, parseErr := parsePkgHelperResponse(out)
if parseErr == nil {
if execErr != nil && resp.OK {
return false, "system_error", "", fmt.Sprintf("pkg-helper fallback exited after successful response: %v", execErr)
}
return resp.OK, resp.Code, resp.Data, resp.Error
}
detail := helperFallbackOutputDetail(out, stderr.String())
if execErr != nil {
return false, "system_error", "", fmt.Sprintf("pkg-helper fallback failed: %v%s", execErr, detail)
}
return false, "helper_error", "", fmt.Sprintf("pkg-helper fallback invalid response: %v%s", parseErr, detail)
}
func helperFallbackOutputDetail(stdout []byte, stderr string) string {
var parts []string
if trimmed := strings.TrimSpace(string(stdout)); trimmed != "" {
parts = append(parts, "stdout: "+trimmed)
}
if trimmed := strings.TrimSpace(stderr); trimmed != "" {
parts = append(parts, "stderr: "+trimmed)
}
if len(parts) == 0 {
return ""
}
return ": " + strings.Join(parts, "; ")
}
func findPkgHelperBinary() (string, bool) {
if path, err := exec.LookPath("pkg-helper"); err == nil {
return path, true
}
return firstExecutableFile(pkgHelperFallbackPaths())
}
func pkgHelperFallbackPaths() []string {
paths := []string{pkgHelperBundledPath}
if exe, err := os.Executable(); err == nil {
paths = append([]string{filepath.Join(filepath.Dir(exe), "pkg-helper")}, paths...)
}
return paths
}
func firstExecutableFile(paths []string) (string, bool) {
seen := make(map[string]struct{}, len(paths))
for _, path := range paths {
if path == "" {
continue
}
if _, ok := seen[path]; ok {
continue
}
seen[path] = struct{}{}
info, err := os.Stat(path)
if err == nil && !info.IsDir() && info.Mode().Perm()&0111 != 0 {
return path, true
}
}
return "", false
} }
// apkViaHelper is the legacy 2-return-value wrapper used by InstallSingleDep, // apkViaHelper is the legacy 2-return-value wrapper used by InstallSingleDep,
+8
View File
@@ -34,6 +34,7 @@ const (
ProviderBytePlus = "byteplus" // BytePlus ModelArk (Seed 2.0 models) ProviderBytePlus = "byteplus" // BytePlus ModelArk (Seed 2.0 models)
ProviderBytePlusCoding = "byteplus_coding" // BytePlus ModelArk Coding Plan ProviderBytePlusCoding = "byteplus_coding" // BytePlus ModelArk Coding Plan
ProviderVertex = "vertex" // Google Cloud Vertex AI (OAuth2 service account + ADC) ProviderVertex = "vertex" // Google Cloud Vertex AI (OAuth2 service account + ADC)
ProviderKimiCoding = "kimi_coding" // Moonshot Kimi Coding (OpenAI-compat, requires fixed User-Agent)
// Novita AI defaults. // Novita AI defaults.
NovitaDefaultAPIBase = "https://api.novita.ai/openai" NovitaDefaultAPIBase = "https://api.novita.ai/openai"
@@ -44,6 +45,12 @@ const (
BytePlusCodingDefaultAPIBase = "https://ark.ap-southeast.bytepluses.com/api/coding/v3" BytePlusCodingDefaultAPIBase = "https://ark.ap-southeast.bytepluses.com/api/coding/v3"
BytePlusDefaultModel = "seed-2-0-lite-260228" BytePlusDefaultModel = "seed-2-0-lite-260228"
// Kimi Coding defaults. The upstream requires a fixed User-Agent on every
// request — handled by the runtime in cmd/gateway_providers.go via
// OpenAIProvider.WithExtraHeaders.
KimiCodingDefaultAPIBase = "https://api.kimi.com/coding/v1"
KimiCodingDefaultModel = "kimi-k2-turbo-preview"
KimiCodingRequiredUserAgent = "claude-code/0.1.0"
) )
// Vertex AI constants live in internal/providers/vertex.go to avoid a store→providers import cycle // Vertex AI constants live in internal/providers/vertex.go to avoid a store→providers import cycle
@@ -77,6 +84,7 @@ var ValidProviderTypes = map[string]bool{
ProviderBytePlus: true, ProviderBytePlus: true,
ProviderBytePlusCoding: true, ProviderBytePlusCoding: true,
ProviderVertex: true, ProviderVertex: true,
ProviderKimiCoding: true,
} }
// VertexProviderSettings holds Vertex-specific config stored in llm_providers.settings JSONB. // VertexProviderSettings holds Vertex-specific config stored in llm_providers.settings JSONB.
+150
View File
@@ -264,3 +264,153 @@ func parseGeminiResponse(respBody []byte) (*providers.ChatResponse, error) {
}, },
}, nil }, nil
} }
// geminiFileUploadStream uploads a file stream to Gemini File API using resumable upload protocol.
// Returns the file name (e.g. "files/abc123") and file URI for use in generateContent.
func geminiFileUploadStream(ctx context.Context, apiKey, displayName string, reader io.Reader, contentLength int64, mime string) (fileName, fileURI string, err error) {
// Step 1: Initiate resumable upload.
initBody, _ := json.Marshal(map[string]any{
"file": map[string]string{"display_name": displayName},
})
initReq, err := http.NewRequestWithContext(ctx, "POST", geminiUploadBase+"?key="+apiKey, bytes.NewReader(initBody))
if err != nil {
return "", "", fmt.Errorf("create init request: %w", err)
}
initReq.Header.Set("Content-Type", "application/json")
initReq.Header.Set("X-Goog-Upload-Protocol", "resumable")
initReq.Header.Set("X-Goog-Upload-Command", "start")
initReq.Header.Set("X-Goog-Upload-Header-Content-Length", fmt.Sprintf("%d", contentLength))
initReq.Header.Set("X-Goog-Upload-Header-Content-Type", mime)
client := &http.Client{Timeout: 60 * time.Second}
initResp, err := client.Do(initReq)
if err != nil {
return "", "", fmt.Errorf("init upload: %w", err)
}
defer initResp.Body.Close()
io.ReadAll(initResp.Body) // drain
if initResp.StatusCode != 200 {
return "", "", fmt.Errorf("init upload HTTP %d", initResp.StatusCode)
}
uploadURL := initResp.Header.Get("X-Goog-Upload-URL")
if uploadURL == "" {
return "", "", fmt.Errorf("no upload URL in response headers")
}
// Step 2: Upload file stream.
uploadReq, err := http.NewRequestWithContext(ctx, "POST", uploadURL, reader)
if err != nil {
return "", "", fmt.Errorf("create upload request: %w", err)
}
uploadReq.ContentLength = contentLength
uploadReq.Header.Set("Content-Length", fmt.Sprintf("%d", contentLength))
uploadReq.Header.Set("X-Goog-Upload-Offset", "0")
uploadReq.Header.Set("X-Goog-Upload-Command", "upload, finalize")
// Do not set global Timeout on the HTTP client here because we are piping a potentially large stream.
uploadClient := &http.Client{}
uploadResp, err := uploadClient.Do(uploadReq)
if err != nil {
return "", "", fmt.Errorf("upload stream: %w", err)
}
defer uploadResp.Body.Close()
respBody, err := io.ReadAll(uploadResp.Body)
if err != nil {
return "", "", fmt.Errorf("read upload response: %w", err)
}
if uploadResp.StatusCode != 200 {
return "", "", fmt.Errorf("upload HTTP %d: %s", uploadResp.StatusCode, truncateStr(string(respBody), 500))
}
var uploadResult struct {
File struct {
Name string `json:"name"`
URI string `json:"uri"`
State string `json:"state"`
} `json:"file"`
}
if err := json.Unmarshal(respBody, &uploadResult); err != nil {
return "", "", fmt.Errorf("parse upload response: %w", err)
}
// Only return URI if file is already ACTIVE; otherwise caller must poll.
if uploadResult.File.State == "ACTIVE" {
return uploadResult.File.Name, uploadResult.File.URI, nil
}
return uploadResult.File.Name, "", nil
}
// geminiFileAPICallStream uploads a file stream via Gemini File API, polls until ready,
// then calls generateContent with file_data reference.
func geminiFileAPICallStream(ctx context.Context, apiKey, model, prompt string, reader io.Reader, contentLength int64, mime string, httpTimeout time.Duration) (*providers.ChatResponse, error) {
displayName := fmt.Sprintf("goclaw_%d", time.Now().UnixNano())
slog.Info("gemini file api: uploading stream", "size", contentLength, "mime", mime)
fileName, fileURI, err := geminiFileUploadStream(ctx, apiKey, displayName, reader, contentLength, mime)
if err != nil {
return nil, fmt.Errorf("upload: %w", err)
}
slog.Info("gemini file api: uploaded stream", "name", fileName)
// If file URI not returned directly, poll for it.
if fileURI == "" {
slog.Info("gemini file api: polling for active state", "name", fileName)
fileURI, err = geminiFilePoll(ctx, apiKey, fileName)
if err != nil {
return nil, fmt.Errorf("poll: %w", err)
}
}
slog.Info("gemini file api: file active", "uri", fileURI)
// Call generateContent with file_data reference.
body := map[string]any{
"contents": []map[string]any{
{
"parts": []map[string]any{
{"file_data": map[string]any{"mime_type": mime, "file_uri": fileURI}},
{"text": prompt},
},
},
},
"generationConfig": map[string]any{
"maxOutputTokens": 16384,
"temperature": 0.2,
},
}
bodyJSON, err := json.Marshal(body)
if err != nil {
return nil, fmt.Errorf("marshal request: %w", err)
}
url := fmt.Sprintf("https://generativelanguage.googleapis.com/v1beta/models/%s:generateContent?key=%s", model, apiKey)
httpReq, err := http.NewRequestWithContext(ctx, "POST", url, bytes.NewReader(bodyJSON))
if err != nil {
return nil, fmt.Errorf("create request: %w", err)
}
httpReq.Header.Set("Content-Type", "application/json")
if httpTimeout == 0 {
httpTimeout = 120 * time.Second
}
client := &http.Client{Timeout: httpTimeout}
httpResp, err := client.Do(httpReq)
if err != nil {
return nil, fmt.Errorf("HTTP request: %w", err)
}
defer httpResp.Body.Close()
respBody, err := io.ReadAll(httpResp.Body)
if err != nil {
return nil, fmt.Errorf("read response: %w", err)
}
if httpResp.StatusCode != 200 {
return nil, fmt.Errorf("HTTP %d: %s", httpResp.StatusCode, truncateStr(string(respBody), 500))
}
// Parse — same response format as geminiNativeDocumentCall.
return parseGeminiResponse(respBody)
}
+40 -3
View File
@@ -10,6 +10,7 @@ import (
"strings" "strings"
"github.com/nextlevelbuilder/goclaw/internal/providers" "github.com/nextlevelbuilder/goclaw/internal/providers"
"github.com/nextlevelbuilder/goclaw/internal/security"
usagecaps "github.com/nextlevelbuilder/goclaw/internal/usage/caps" usagecaps "github.com/nextlevelbuilder/goclaw/internal/usage/caps"
) )
@@ -62,7 +63,7 @@ func (t *ReadImageTool) SetUsageCapService(svc *usagecaps.Service) {
func (t *ReadImageTool) Name() string { return "read_image" } func (t *ReadImageTool) Name() string { return "read_image" }
func (t *ReadImageTool) Description() string { func (t *ReadImageTool) Description() string {
return "Analyze images using vision AI. Works with: (1) images sent by the user (<media:image> tags), (2) workspace/generated image files (pass a file path)." return "Analyze images using vision AI. Works with images sent by the user, workspace/generated image files, or public HTTP/HTTPS image URLs."
} }
func (t *ReadImageTool) Parameters() map[string]any { func (t *ReadImageTool) Parameters() map[string]any {
@@ -77,6 +78,10 @@ func (t *ReadImageTool) Parameters() map[string]any {
"type": "string", "type": "string",
"description": "Optional file path to an image in the workspace. Use this for generated images or attachments. If omitted, analyzes images from the conversation.", "description": "Optional file path to an image in the workspace. Use this for generated images or attachments. If omitted, analyzes images from the conversation.",
}, },
"url": map[string]any{
"type": "string",
"description": "Optional URL to an image. Use this to analyze images hosted online.",
},
}, },
"required": []string{"prompt"}, "required": []string{"prompt"},
} }
@@ -91,18 +96,32 @@ func (t *ReadImageTool) Execute(ctx context.Context, args map[string]any) *Resul
prompt = "Describe this image in detail." prompt = "Describe this image in detail."
} }
imgPath, _ := args["path"].(string)
imgURL, _ := args["url"].(string)
if imgPath != "" && imgURL != "" {
return ErrorResult("Both 'path' and 'url' parameters cannot be specified. Choose only one.")
}
// If path is provided, load image from workspace file // If path is provided, load image from workspace file
images := MediaImagesFromCtx(ctx) images := MediaImagesFromCtx(ctx)
if imgPath, _ := args["path"].(string); imgPath != "" { if imgPath != "" {
fileImages, err := t.loadImageFromPath(ctx, imgPath) fileImages, err := t.loadImageFromPath(ctx, imgPath)
if err != nil { if err != nil {
return ErrorResult(err.Error()) return ErrorResult(err.Error())
} }
images = fileImages images = fileImages
} else if imgURL != "" {
if _, _, err := security.Validate(imgURL); err != nil {
return ErrorResult(fmt.Sprintf("Invalid image URL: %v", err))
}
images = []providers.ImageContent{{
URL: imgURL,
}}
} }
if len(images) == 0 { if len(images) == 0 {
return ErrorResult("No images available. Either send an image in the chat or provide a file path with the 'path' parameter.") return ErrorResult("No images available. Either send an image in the chat, provide a file path with 'path', or provide an image URL with 'url'.")
} }
chain := ResolveMediaProviderChain(ctx, "read_image", "", "", chain := ResolveMediaProviderChain(ctx, "read_image", "", "",
@@ -138,6 +157,24 @@ func (t *ReadImageTool) callProvider(ctx context.Context, cp credentialProvider,
prompt := GetParamString(params, "prompt", "Describe this image in detail.") prompt := GetParamString(params, "prompt", "Describe this image in detail.")
images, _ := params["images"].([]providers.ImageContent) images, _ := params["images"].([]providers.ImageContent)
for _, img := range images {
if img.URL == "" {
continue
}
if _, _, err := security.Validate(img.URL); err != nil {
return nil, nil, fmt.Errorf("invalid image URL: %w", err)
}
}
// Anthropic Claude does not support URL references and requires base64-encoded image data.
if providerName == "anthropic" || providerName == "claude-cli" {
for _, img := range images {
if img.URL != "" && img.Data == "" {
return nil, nil, fmt.Errorf("provider %q does not support analyzing images directly from a URL", providerName)
}
}
}
// Get the full provider for Chat() access // Get the full provider for Chat() access
p, err := t.registry.Get(ctx, providerName) p, err := t.registry.Get(ctx, providerName)
if err != nil { if err != nil {
+71
View File
@@ -0,0 +1,71 @@
package tools
import (
"context"
"strings"
"testing"
"github.com/nextlevelbuilder/goclaw/internal/providers"
)
func TestReadImage_BothPathAndUrl_Error(t *testing.T) {
tool := NewReadImageTool(nil)
res := tool.Execute(context.Background(), map[string]any{
"prompt": "describe this",
"path": "workspace/image.png",
"url": "https://example.com/image.png",
})
if !res.IsError {
t.Fatalf("expected error when both path and url are provided")
}
if !strings.Contains(res.ForLLM, "Both 'path' and 'url' parameters cannot be specified") {
t.Errorf("unexpected error message: %s", res.ForLLM)
}
}
func TestReadImage_PrivateURL_Error(t *testing.T) {
tool := NewReadImageTool(nil)
res := tool.Execute(context.Background(), map[string]any{
"prompt": "describe this",
"url": "http://127.0.0.1/image.png",
})
if !res.IsError {
t.Fatalf("expected error for private image URL")
}
if !strings.Contains(res.ForLLM, "Invalid image URL") {
t.Errorf("unexpected error message: %s", res.ForLLM)
}
}
func TestReadImage_AnthropicURL_Error(t *testing.T) {
tool := NewReadImageTool(nil)
params := map[string]any{
"prompt": "describe this",
"images": []providers.ImageContent{
{
URL: "https://93.184.216.34/image.png",
},
},
}
_, _, err := tool.callProvider(context.Background(), nil, "anthropic", "claude-3-sonnet", params)
if err == nil {
t.Fatalf("expected error for anthropic provider with image URL")
}
if !strings.Contains(err.Error(), "does not support analyzing images directly from a URL") {
t.Errorf("unexpected error message: %v", err)
}
// Should also error for claude-cli
_, _, err = tool.callProvider(context.Background(), nil, "claude-cli", "claude-3-sonnet", params)
if err == nil {
t.Fatalf("expected error for claude-cli provider with image URL")
}
}
+46 -11
View File
@@ -4,10 +4,13 @@ import (
"context" "context"
"fmt" "fmt"
"log/slog" "log/slog"
"net"
"os" "os"
"path/filepath"
"strings" "strings"
"github.com/nextlevelbuilder/goclaw/internal/providers" "github.com/nextlevelbuilder/goclaw/internal/providers"
"github.com/nextlevelbuilder/goclaw/internal/security"
usagecaps "github.com/nextlevelbuilder/goclaw/internal/usage/caps" usagecaps "github.com/nextlevelbuilder/goclaw/internal/usage/caps"
) )
@@ -31,6 +34,8 @@ func MediaVideoRefsFromCtx(ctx context.Context) []providers.MediaRef {
// videoMaxBytes is the max file size for video analysis (100MB). // videoMaxBytes is the max file size for video analysis (100MB).
const videoMaxBytes = 100 * 1024 * 1024 const videoMaxBytes = 100 * 1024 * 1024
const videoURLPinnedIPParam = "_pinned_ip"
// videoProviderPriority is the order in which providers are tried for video analysis. // videoProviderPriority is the order in which providers are tried for video analysis.
// OpenAI excluded — no native video upload in chat completions. // OpenAI excluded — no native video upload in chat completions.
var videoProviderPriority = []string{"gemini", "openrouter"} var videoProviderPriority = []string{"gemini", "openrouter"}
@@ -77,6 +82,10 @@ func (t *ReadVideoTool) Parameters() map[string]any {
"type": "string", "type": "string",
"description": "Optional: specific media_id from <media:video> tag. If omitted, uses most recent video.", "description": "Optional: specific media_id from <media:video> tag. If omitted, uses most recent video.",
}, },
"url": map[string]any{
"type": "string",
"description": "Optional URL to a video file. Use this to analyze videos hosted online.",
},
}, },
"required": []string{"prompt"}, "required": []string{"prompt"},
} }
@@ -88,21 +97,43 @@ func (t *ReadVideoTool) Execute(ctx context.Context, args map[string]any) *Resul
prompt = "Analyze this video and describe its contents." prompt = "Analyze this video and describe its contents."
} }
mediaID, _ := args["media_id"].(string) mediaID, _ := args["media_id"].(string)
videoURL, _ := args["url"].(string)
videoPath, videoMime, err := t.resolveVideoFile(ctx, mediaID) if mediaID != "" && videoURL != "" {
if err != nil { return ErrorResult("Both 'media_id' and 'url' parameters cannot be specified. Choose only one.")
return ErrorResult(err.Error())
} }
slog.Info("read_video: resolved file", "path", videoPath, "mime", videoMime, "media_id", mediaID) var data []byte
var videoMime string
var pinnedIP net.IP
data, err := os.ReadFile(videoPath) if videoURL != "" {
if err != nil { validatedURL, validatedIP, err := security.Validate(videoURL)
return ErrorResult(fmt.Sprintf("Failed to read video file: %v", err)) if err != nil {
} return ErrorResult(fmt.Sprintf("Invalid video URL: %v", err))
slog.Info("read_video: file loaded", "size_bytes", len(data)) }
if len(data) > videoMaxBytes { pinnedIP = validatedIP
return ErrorResult(fmt.Sprintf("Video too large: %d bytes (max %d)", len(data), videoMaxBytes))
// Infer MIME type from URL extension
ext := filepath.Ext(validatedURL.Path)
videoMime = mimeFromVideoExt(ext)
} else {
videoPath, mime, err := t.resolveVideoFile(ctx, mediaID)
if err != nil {
return ErrorResult(err.Error())
}
videoMime = mime
slog.Info("read_video: resolved file", "path", videoPath, "mime", videoMime, "media_id", mediaID)
fileData, err := os.ReadFile(videoPath)
if err != nil {
return ErrorResult(fmt.Sprintf("Failed to read video file: %v", err))
}
slog.Info("read_video: file loaded", "size_bytes", len(fileData))
if len(fileData) > videoMaxBytes {
return ErrorResult(fmt.Sprintf("Video too large: %d bytes (max %d)", len(fileData), videoMaxBytes))
}
data = fileData
} }
chain := ResolveMediaProviderChain(ctx, "read_video", "", "", chain := ResolveMediaProviderChain(ctx, "read_video", "", "",
@@ -114,7 +145,11 @@ func (t *ReadVideoTool) Execute(ctx context.Context, args map[string]any) *Resul
} }
chain[i].Params["prompt"] = prompt chain[i].Params["prompt"] = prompt
chain[i].Params["data"] = data chain[i].Params["data"] = data
chain[i].Params["url"] = videoURL
chain[i].Params["mime"] = videoMime chain[i].Params["mime"] = videoMime
if pinnedIP != nil {
chain[i].Params[videoURLPinnedIPParam] = pinnedIP
}
} }
chainResult, err := ExecuteWithChain(ctx, chain, t.registry, t.callProvider) chainResult, err := ExecuteWithChain(ctx, chain, t.registry, t.callProvider)
+92 -6
View File
@@ -5,10 +5,14 @@ import (
"encoding/base64" "encoding/base64"
"fmt" "fmt"
"log/slog" "log/slog"
"net"
"net/http"
"path/filepath" "path/filepath"
"strings"
"time" "time"
"github.com/nextlevelbuilder/goclaw/internal/providers" "github.com/nextlevelbuilder/goclaw/internal/providers"
"github.com/nextlevelbuilder/goclaw/internal/security"
) )
// resolveVideoFile finds the video file path from context MediaRefs. // resolveVideoFile finds the video file path from context MediaRefs.
@@ -60,16 +64,26 @@ func (t *ReadVideoTool) resolveVideoFile(ctx context.Context, mediaID string) (p
// callProvider dispatches video analysis to the appropriate provider API. // callProvider dispatches video analysis to the appropriate provider API.
// Gemini: uses File API (upload → poll → file_data in generateContent). // Gemini: uses File API (upload → poll → file_data in generateContent).
// Others: falls back to base64 in image_url (OpenRouter routes to Gemini which handles video). // Others: falls back to base64 or URL in video_url (OpenRouter routes to Gemini which handles video).
func (t *ReadVideoTool) callProvider(ctx context.Context, cp credentialProvider, providerName, model string, params map[string]any) ([]byte, *providers.Usage, error) { func (t *ReadVideoTool) callProvider(ctx context.Context, cp credentialProvider, providerName, model string, params map[string]any) ([]byte, *providers.Usage, error) {
prompt := GetParamString(params, "prompt", "Analyze this video and describe its contents.") prompt := GetParamString(params, "prompt", "Analyze this video and describe its contents.")
data, _ := params["data"].([]byte) data, _ := params["data"].([]byte)
videoURL, _ := params["url"].(string)
mime := GetParamString(params, "mime", "video/mp4") mime := GetParamString(params, "mime", "video/mp4")
pinnedIP, _ := params[videoURLPinnedIPParam].(net.IP)
if videoURL != "" && pinnedIP == nil {
var err error
if _, pinnedIP, err = security.Validate(videoURL); err != nil {
return nil, nil, fmt.Errorf("invalid video URL: %w", err)
}
}
// Gemini: use File API (requires credentials). // Gemini: use File API (requires credentials).
ptype := GetParamString(params, "_provider_type", providerTypeFromName(providerName)) ptype := GetParamString(params, "_provider_type", providerTypeFromName(providerName))
if cp != nil && ptype == "gemini" { if cp != nil && ptype == "gemini" {
slog.Info("read_video: using gemini file API", "provider", providerName, "model", model, "size", len(data), "mime", mime) var resp *providers.ChatResponse
var err error
chatReq := providers.ChatRequest{ chatReq := providers.ChatRequest{
Messages: []providers.Message{{Role: "user", Content: prompt}}, Messages: []providers.Message{{Role: "user", Content: prompt}},
Model: model, Model: model,
@@ -79,7 +93,71 @@ func (t *ReadVideoTool) callProvider(ctx context.Context, cp credentialProvider,
if reserveErr != nil { if reserveErr != nil {
return nil, nil, reserveErr return nil, nil, reserveErr
} }
resp, err := geminiFileAPICall(ctx, cp.APIKey(), model, prompt, data, mime, 180*time.Second)
if videoURL != "" {
slog.Info("read_video: streaming URL directly to Gemini File API", "provider", providerName, "model", model, "url", videoURL)
// Send GET request to fetch the stream.
reqCtx := security.WithPinnedIP(ctx, pinnedIP)
req, getErr := http.NewRequestWithContext(reqCtx, "GET", videoURL, nil)
if getErr != nil {
if reservation != nil {
reservation.Reconcile(ctx, nil, getErr)
}
return nil, nil, fmt.Errorf("failed to create GET request for video URL: %w", getErr)
}
// Use the shared SSRF-safe client so DNS stays pinned during streaming.
client := security.NewSafeClient(0)
httpResp, getErr := client.Do(req)
if getErr != nil {
if reservation != nil {
reservation.Reconcile(ctx, nil, getErr)
}
return nil, nil, fmt.Errorf("failed to fetch video URL: %w", getErr)
}
defer httpResp.Body.Close()
if httpResp.StatusCode < 200 || httpResp.StatusCode >= 300 {
statusErr := fmt.Errorf("video URL returned status code %d", httpResp.StatusCode)
if reservation != nil {
reservation.Reconcile(ctx, nil, statusErr)
}
return nil, nil, statusErr
}
// Validate Content-Length
contentLength := httpResp.ContentLength
if contentLength <= 0 {
invalidLenErr := fmt.Errorf("URL does not support static streaming (missing or invalid Content-Length: %d)", contentLength)
if reservation != nil {
reservation.Reconcile(ctx, nil, invalidLenErr)
}
return nil, nil, invalidLenErr
}
// Check limits: maximum 2 GB
const videoMaxStreamBytes = 2 * 1024 * 1024 * 1024 // 2 GB
if contentLength > videoMaxStreamBytes {
limitErr := fmt.Errorf("video stream size (%d bytes) exceeds the maximum limit of 2 GB", contentLength)
if reservation != nil {
reservation.Reconcile(ctx, nil, limitErr)
}
return nil, nil, limitErr
}
// Extract MIME type from Content-Type header if valid video type, otherwise use the inferred/passed one
contentType := httpResp.Header.Get("Content-Type")
if contentType != "" && strings.HasPrefix(contentType, "video/") {
mime = contentType
}
resp, err = geminiFileAPICallStream(ctx, cp.APIKey(), model, prompt, httpResp.Body, contentLength, mime, 300*time.Second)
} else {
slog.Info("read_video: using gemini file API", "provider", providerName, "model", model, "size", len(data), "mime", mime)
resp, err = geminiFileAPICall(ctx, cp.APIKey(), model, prompt, data, mime, 180*time.Second)
}
if reservation != nil { if reservation != nil {
reservation.Reconcile(ctx, resp, err) reservation.Reconcile(ctx, resp, err)
} }
@@ -89,19 +167,27 @@ func (t *ReadVideoTool) callProvider(ctx context.Context, cp credentialProvider,
return []byte(resp.Content), resp.Usage, nil return []byte(resp.Content), resp.Usage, nil
} }
// Other providers: try standard Chat API with base64 as image_url (best effort). // Other providers: try standard Chat API with base64 or URL as video_url (best effort).
p, err := t.registry.Get(ctx, providerName) p, err := t.registry.Get(ctx, providerName)
if err != nil { if err != nil {
return nil, nil, fmt.Errorf("provider %q not available: %w", providerName, err) return nil, nil, fmt.Errorf("provider %q not available: %w", providerName, err)
} }
slog.Info("read_video: using chat API fallback", "provider", providerName, "model", model, "size", len(data)) var vidContent providers.VideoContent
if videoURL != "" {
slog.Info("read_video: using chat API with direct video URL", "provider", providerName, "model", model, "url", videoURL)
vidContent = providers.VideoContent{MimeType: mime, URL: videoURL}
} else {
slog.Info("read_video: using chat API fallback with base64", "provider", providerName, "model", model, "size", len(data))
vidContent = providers.VideoContent{MimeType: mime, Data: base64.StdEncoding.EncodeToString(data)}
}
chatReq := providers.ChatRequest{ chatReq := providers.ChatRequest{
Messages: []providers.Message{ Messages: []providers.Message{
{ {
Role: "user", Role: "user",
Content: prompt, Content: prompt,
Images: []providers.ImageContent{{MimeType: mime, Data: base64.StdEncoding.EncodeToString(data)}}, Videos: []providers.VideoContent{vidContent},
}, },
}, },
Model: model, Model: model,
+123
View File
@@ -0,0 +1,123 @@
package tools
import (
"context"
"net/http"
"net/http/httptest"
"strings"
"testing"
"github.com/nextlevelbuilder/goclaw/internal/security"
)
type mockCredentialProvider struct {
apiKey string
apiBase string
}
func (m *mockCredentialProvider) APIKey() string { return m.apiKey }
func (m *mockCredentialProvider) APIBase() string { return m.apiBase }
func TestReadVideo_BothMediaIdAndUrl_Error(t *testing.T) {
tool := NewReadVideoTool(nil, nil)
res := tool.Execute(context.Background(), map[string]any{
"prompt": "describe this video",
"media_id": "video-123",
"url": "https://example.com/video.mp4",
})
if !res.IsError {
t.Fatalf("expected error when both media_id and url are provided")
}
if !strings.Contains(res.ForLLM, "Both 'media_id' and 'url' parameters cannot be specified") {
t.Errorf("unexpected error message: %s", res.ForLLM)
}
}
func TestReadVideo_PrivateURL_Error(t *testing.T) {
tool := NewReadVideoTool(nil, nil)
res := tool.Execute(context.Background(), map[string]any{
"prompt": "describe this video",
"url": "http://127.0.0.1/video.mp4",
})
if !res.IsError {
t.Fatalf("expected error for private video URL")
}
if !strings.Contains(res.ForLLM, "Invalid video URL") {
t.Errorf("unexpected error message: %s", res.ForLLM)
}
}
func TestReadVideo_GeminiURL_Validation(t *testing.T) {
security.SetAllowLoopbackForTest(true)
defer security.SetAllowLoopbackForTest(false)
// Missing Content-Length should fail before upload.
ts1 := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
w.Header().Set("Transfer-Encoding", "chunked")
w.Write([]byte("chunked data mock video"))
}))
defer ts1.Close()
tool := NewReadVideoTool(nil, nil)
cp := &mockCredentialProvider{apiKey: "test-key"}
params1 := map[string]any{
"prompt": "describe this video",
"url": ts1.URL,
"_provider_type": "gemini",
}
_, _, err := tool.callProvider(context.Background(), cp, "gemini", "gemini-2.5-flash", params1)
if err == nil {
t.Fatalf("expected error for missing Content-Length")
}
if !strings.Contains(err.Error(), "URL does not support static streaming") {
t.Errorf("unexpected error for missing Content-Length: %v", err)
}
// Content-Length over the Gemini File API limit should fail before upload.
ts2 := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
w.Header().Set("Content-Length", "2147483649") // 2GB + 1 byte
w.WriteHeader(http.StatusOK)
}))
defer ts2.Close()
params2 := map[string]any{
"prompt": "describe this video",
"url": ts2.URL,
"_provider_type": "gemini",
}
_, _, err = tool.callProvider(context.Background(), cp, "gemini", "gemini-2.5-flash", params2)
if err == nil {
t.Fatalf("expected error for Content-Length exceeding 2GB")
}
if !strings.Contains(err.Error(), "exceeds the maximum limit of 2 GB") {
t.Errorf("unexpected error for limit exceed: %v", err)
}
// Non-2xx status should be reported before upload.
ts3 := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
w.WriteHeader(http.StatusNotFound)
}))
defer ts3.Close()
params3 := map[string]any{
"prompt": "describe this video",
"url": ts3.URL,
"_provider_type": "gemini",
}
_, _, err = tool.callProvider(context.Background(), cp, "gemini", "gemini-2.5-flash", params3)
if err == nil {
t.Fatalf("expected error for HTTP 404 status code")
}
if !strings.Contains(err.Error(), "video URL returned status code 404") {
t.Errorf("unexpected error for HTTP status code: %v", err)
}
}
+1 -1
View File
@@ -254,7 +254,7 @@ func resolveRemoteCDP(remoteURL string) (string, error) {
defer resp.Body.Close() defer resp.Body.Close()
if resp.StatusCode != http.StatusOK { if resp.StatusCode != http.StatusOK {
return "", fmt.Errorf("/json/version returned HTTP %d", resp.StatusCode) return "", fmt.Errorf("query /json/version at %s returned HTTP %d", versionURL, resp.StatusCode)
} }
var ver struct { var ver struct {
+1
View File
@@ -33,6 +33,7 @@ export const PROVIDER_TYPES: ProviderTypeInfo[] = [
{ value: "zai_coding", label: "Z.ai Coding Plan", apiBase: "https://api.z.ai/api/coding/paas/v4", placeholder: "" }, { value: "zai_coding", label: "Z.ai Coding Plan", apiBase: "https://api.z.ai/api/coding/paas/v4", placeholder: "" },
{ value: "byteplus", label: "BytePlus ModelArk", apiBase: "https://ark.ap-southeast.bytepluses.com/api/v3", placeholder: "" }, { value: "byteplus", label: "BytePlus ModelArk", apiBase: "https://ark.ap-southeast.bytepluses.com/api/v3", placeholder: "" },
{ value: "byteplus_coding", label: "BytePlus Coding Plan", apiBase: "https://ark.ap-southeast.bytepluses.com/api/coding/v3", placeholder: "" }, { value: "byteplus_coding", label: "BytePlus Coding Plan", apiBase: "https://ark.ap-southeast.bytepluses.com/api/coding/v3", placeholder: "" },
{ value: "kimi_coding", label: "Kimi Coding (Moonshot)", apiBase: "https://api.kimi.com/coding/v1", placeholder: "" },
{ value: "ollama", label: "Ollama (Local)", apiBase: "http://localhost:11434/v1", placeholder: "" }, { value: "ollama", label: "Ollama (Local)", apiBase: "http://localhost:11434/v1", placeholder: "" },
{ value: "ollama_cloud", label: "Ollama Cloud", apiBase: "https://ollama.com/v1", placeholder: "" }, { value: "ollama_cloud", label: "Ollama Cloud", apiBase: "https://ollama.com/v1", placeholder: "" },
{ value: "claude_cli", label: "Claude CLI (Local)", apiBase: "", placeholder: "" }, { value: "claude_cli", label: "Claude CLI (Local)", apiBase: "", placeholder: "" },
+1 -1
View File
@@ -721,7 +721,7 @@
}, },
"errors": { "errors": {
"invalidName": "Use lowercase letters, digits, hyphens, underscores (2-64 chars).", "invalidName": "Use lowercase letters, digits, hyphens, underscores (2-64 chars).",
"invalidDomain": "Must be a valid Bitrix24 portal domain (e.g. mycorp.bitrix24.com).", "invalidDomain": "Must be a valid hostname (e.g. mycorp.bitrix24.com, mycorp.bitrix24.vn, or your self-hosted domain).",
"duplicateName": "A portal with this name already exists.", "duplicateName": "A portal with this name already exists.",
"forbidden": "You need tenant admin permission to create portals.", "forbidden": "You need tenant admin permission to create portals.",
"gatewayURLUnknown": "Open the goclaw UI via your public URL first (not localhost), then retry." "gatewayURLUnknown": "Open the goclaw UI via your public URL first (not localhost), then retry."
+1 -1
View File
@@ -639,7 +639,7 @@
}, },
"errors": { "errors": {
"invalidName": "Chữ thường, số, gạch nối, gạch dưới (2-64 ký tự).", "invalidName": "Chữ thường, số, gạch nối, gạch dưới (2-64 ký tự).",
"invalidDomain": "Phải là domain Bitrix24 hợp lệ (vd: mycorp.bitrix24.com).", "invalidDomain": "Phải là hostname hợp lệ (vd: mycorp.bitrix24.com, mycorp.bitrix24.vn, hoặc domain self-hosted của bạn).",
"duplicateName": "Tên portal đã tồn tại.", "duplicateName": "Tên portal đã tồn tại.",
"forbidden": "Cần quyền tenant admin để tạo portal.", "forbidden": "Cần quyền tenant admin để tạo portal.",
"gatewayURLUnknown": "Mở UI goclaw qua public URL (không phải localhost) trước, rồi retry." "gatewayURLUnknown": "Mở UI goclaw qua public URL (không phải localhost) trước, rồi retry."
+1 -1
View File
@@ -639,7 +639,7 @@
}, },
"errors": { "errors": {
"invalidName": "请使用小写字母、数字、连字符、下划线(2-64 字符)。", "invalidName": "请使用小写字母、数字、连字符、下划线(2-64 字符)。",
"invalidDomain": "必须是有效的 Bitrix24 域名(例如:mycorp.bitrix24.com)。", "invalidDomain": "必须是有效的主机名(例如:mycorp.bitrix24.com、mycorp.bitrix24.vn 或你的自托管域名)。",
"duplicateName": "门户名称已存在。", "duplicateName": "门户名称已存在。",
"forbidden": "您需要租户管理员权限才能创建门户。", "forbidden": "您需要租户管理员权限才能创建门户。",
"gatewayURLUnknown": "请先通过公网 URL(非 localhost)打开 goclaw UI 然后重试。" "gatewayURLUnknown": "请先通过公网 URL(非 localhost)打开 goclaw UI 然后重试。"
@@ -10,13 +10,72 @@ import { useBitrixPortalCreate } from "./use-bitrix-portals";
// Validation mirrors the server-side regex in // Validation mirrors the server-side regex in
// internal/gateway/methods/bitrix_portals.go. Server is authoritative; // internal/gateway/methods/bitrix_portals.go. Server is authoritative;
// client validation is purely UX so the operator gets feedback before a // client validation is purely UX so the operator gets feedback before a
// round-trip. Pattern intentionally accepts a wide TLD set (Bitrix24 has // round-trip. Pattern accepts Bitrix24 regional clouds (.com, .eu, .ru,
// regional clouds: .com, .eu, .ru, .de, .fr, .jp, .in, .kz, .ua, .by) plus // .de, .fr, .jp, .in, .kz, .ua, .by, .vn, .tr, .es, .com.br, .com.ar),
// .bitrix.info for self-hosted. // .bitrix.info self-hosted, plus any valid hostname/port for fully
// self-hosted custom domains.
const BITRIX_DOMAIN_RE = const BITRIX_DOMAIN_RE =
/^[a-z0-9]([a-z0-9-]{0,61}[a-z0-9])?\.(bitrix24\.(com|eu|ru|de|fr|jp|in|kz|ua|by)|bitrix\.info)$/; /^[a-z0-9]([a-z0-9-]{0,61}[a-z0-9])?\.(bitrix24\.(com|eu|ru|de|fr|jp|in|kz|ua|by|vn|tr|es|com\.br|com\.ar)|bitrix\.info)$/;
const SELF_HOSTED_DOMAIN_RE =
/^[a-z0-9]([a-z0-9-]{0,61}[a-z0-9])?(\.[a-z0-9]([a-z0-9-]{0,61}[a-z0-9])?)*(:\d+)?$/;
const PORTAL_NAME_RE = /^[a-z0-9][a-z0-9_-]{0,62}[a-z0-9]$/; const PORTAL_NAME_RE = /^[a-z0-9][a-z0-9_-]{0,62}[a-z0-9]$/;
// validateSelfHostedDomain mirrors the backend SSRF + port validation.
// Rejects localhost, .local, .localhost TLDs, literal private/loopback IPs,
// and invalid port ranges (0, >65535).
function validateSelfHostedDomain(domain: string): string | null {
// Extract host and optional port.
let host = domain;
let portStr: string | undefined;
const colonIdx = domain.lastIndexOf(":");
if (colonIdx !== -1) {
host = domain.slice(0, colonIdx);
portStr = domain.slice(colonIdx + 1);
}
// Validate port range.
if (portStr !== undefined) {
const port = Number(portStr);
if (!Number.isInteger(port) || port < 1 || port > 65535) {
return "port must be 1-65535";
}
}
// Reject localhost and .local/.localhost TLDs.
const lowerHost = host.toLowerCase();
if (
lowerHost === "localhost" ||
lowerHost.endsWith(".localhost") ||
lowerHost.endsWith(".local")
) {
return "private/internal hostnames (localhost, .local, .localhost) are not allowed";
}
// Reject literal private/loopback IPs.
// Simple check for common patterns — backend does full CIDR validation.
if (
lowerHost === "127.0.0.1" ||
lowerHost.startsWith("127.") ||
lowerHost === "10.0.0.0" ||
lowerHost.startsWith("10.") ||
lowerHost.startsWith("192.168.") ||
lowerHost.startsWith("172.16.") ||
lowerHost.startsWith("172.17.") ||
lowerHost.startsWith("172.18.") ||
lowerHost.startsWith("172.19.") ||
lowerHost.startsWith("172.2") ||
lowerHost.startsWith("172.3") ||
lowerHost === "169.254.169.254" ||
lowerHost.startsWith("169.254.") ||
lowerHost === "::1" ||
lowerHost === "0.0.0.0"
) {
return "IP is in a blocked range (loopback/private/metadata)";
}
return null;
}
interface BitrixPortalFormStepProps { interface BitrixPortalFormStepProps {
/** Invoked with the server response after bitrix.portals.create succeeds. */ /** Invoked with the server response after bitrix.portals.create succeeds. */
onSuccess: (createdName: string, installUrl: string, warning?: string) => void; onSuccess: (createdName: string, installUrl: string, warning?: string) => void;
@@ -53,10 +112,21 @@ export function BitrixPortalFormStep({ onSuccess, onCancel }: BitrixPortalFormSt
defaultValue: "Use lowercase letters, digits, hyphens, underscores (2-64 chars).", defaultValue: "Use lowercase letters, digits, hyphens, underscores (2-64 chars).",
}); });
} }
if (!BITRIX_DOMAIN_RE.test(domain.toLowerCase())) { const domainLower = domain.toLowerCase();
const isCloud = BITRIX_DOMAIN_RE.test(domainLower);
const isSelfHostedSyntax = SELF_HOSTED_DOMAIN_RE.test(domainLower);
if (!isCloud && !isSelfHostedSyntax) {
e.domain = t("bitrix24.create.errors.invalidDomain", { e.domain = t("bitrix24.create.errors.invalidDomain", {
defaultValue: "Must be a valid Bitrix24 portal domain (e.g. mycorp.bitrix24.com).", defaultValue: "Must be a valid hostname (e.g. *.bitrix24.com, *.bitrix.info, or your self-hosted domain).",
}); });
} else if (!isCloud) {
// SSRF + port validation for self-hosted domains.
const ssrfErr = validateSelfHostedDomain(domainLower);
if (ssrfErr) {
e.domain = t("bitrix24.create.errors.invalidDomain", {
defaultValue: ssrfErr,
});
}
} }
if (!clientId.trim()) e.client_id = t("common.required", { defaultValue: "Required" }); if (!clientId.trim()) e.client_id = t("common.required", { defaultValue: "Required" });
if (!clientSecret.trim()) e.client_secret = t("common.required", { defaultValue: "Required" }); if (!clientSecret.trim()) e.client_secret = t("common.required", { defaultValue: "Required" });
@@ -121,7 +191,7 @@ export function BitrixPortalFormStep({ onSuccess, onCancel }: BitrixPortalFormSt
value={domain} value={domain}
onChange={(e) => setDomain(e.target.value)} onChange={(e) => setDomain(e.target.value)}
onBlur={handleDomainBlur} onBlur={handleDomainBlur}
placeholder="tamgiac.bitrix24.com" placeholder="mycorp.bitrix24.vn or bitrix.example.com"
autoComplete="off" autoComplete="off"
autoFocus autoFocus
/> />
+2 -2
View File
@@ -4,8 +4,8 @@ export function SetupLayout({ children }: { children: React.ReactNode }) {
const { t } = useTranslation("setup"); const { t } = useTranslation("setup");
return ( return (
<div className="flex min-h-dvh items-center justify-center bg-background px-4 py-8"> <div className="flex min-h-dvh items-start justify-center bg-background px-4 py-8 sm:items-center">
<div className="w-full max-w-2xl space-y-6"> <div className="w-full max-w-2xl space-y-6 overflow-y-auto max-h-dvh sm:max-h-none">
<div className="text-center"> <div className="text-center">
<img src="/goclaw-icon.svg" alt="GoClaw" className="mx-auto mb-4 h-16 w-16" /> <img src="/goclaw-icon.svg" alt="GoClaw" className="mx-auto mb-4 h-16 w-16" />
<h1 className="text-4xl font-bold tracking-tight">GoClaw Setup</h1> <h1 className="text-4xl font-bold tracking-tight">GoClaw Setup</h1>