mirror of
https://github.com/tiennm99/goclaw.git
synced 2026-10-11 03:13:24 +00:00
Merge remote-tracking branch 'upstream/dev' into dev
# Conflicts: # internal/cron/service.go
This commit is contained in:
commit
591d809779
39 files changed
+1705
-92
No files matched your search
@@ -397,6 +397,19 @@ func registerProvidersFromDB(registry *providers.Registry, provStore store.Provi
|
||||
prov := providers.NewOpenAIProvider(p.Name, p.APIKey, base, store.BytePlusDefaultModel)
|
||||
prov.WithProviderType(p.ProviderType)
|
||||
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:
|
||||
prov := providers.NewOpenAIProvider(p.Name, p.APIKey, p.APIBase, "")
|
||||
prov.WithProviderType(p.ProviderType)
|
||||
|
||||
@@ -131,6 +131,7 @@ func seedConfigForContext(ctx context.Context, sc store.SystemConfigStore, cfg *
|
||||
set("tts.auto", cfg.Tts.Auto)
|
||||
set("tts.mode", cfg.Tts.Mode)
|
||||
setInt("tts.max_length", cfg.Tts.MaxLength)
|
||||
setInt("tts.timeout_ms", cfg.Tts.TimeoutMs)
|
||||
|
||||
// Cron
|
||||
setInt("cron.max_retries", cfg.Cron.MaxRetries)
|
||||
|
||||
+32
-10
@@ -60,6 +60,20 @@ type response struct {
|
||||
}
|
||||
|
||||
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")
|
||||
|
||||
// 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 {
|
||||
apkMutex.Lock()
|
||||
defer apkMutex.Unlock()
|
||||
|
||||
slog.Info("pkg-helper: installing", "package", pkg)
|
||||
|
||||
cmd := exec.Command("apk", "add", "--no-cache", pkg)
|
||||
out, err := cmd.CombinedOutput()
|
||||
out, err := runApkFunc("add", "--no-cache", pkg)
|
||||
if err != nil {
|
||||
msg, code := classifyApkOutput(string(out), err)
|
||||
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)
|
||||
|
||||
cmd := exec.Command("apk", "del", pkg)
|
||||
out, err := cmd.CombinedOutput()
|
||||
out, err := runApkFunc("del", pkg)
|
||||
if err != nil {
|
||||
msg, code := classifyApkOutput(string(out), err)
|
||||
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)
|
||||
|
||||
cmd := exec.Command("apk", "add", "-u", pkg)
|
||||
out, err := cmd.CombinedOutput()
|
||||
out, err := runApkFunc("add", "-u", pkg)
|
||||
if err != nil {
|
||||
msg, code := classifyApkOutput(string(out), err)
|
||||
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")
|
||||
|
||||
cmd := exec.Command("apk", "update")
|
||||
out, err := cmd.CombinedOutput()
|
||||
out, err := runApkFunc("update")
|
||||
if err != nil {
|
||||
msg, code := classifyApkOutput(string(out), err)
|
||||
slog.Warn("pkg-helper: update-index failed", "error", msg, "code", code)
|
||||
@@ -292,8 +315,7 @@ func doListOutdated() response {
|
||||
apkMutex.Lock()
|
||||
defer apkMutex.Unlock()
|
||||
|
||||
cmd := exec.Command("apk", "version", "-l", "<")
|
||||
out, err := cmd.CombinedOutput()
|
||||
out, err := runApkFunc("version", "-l", "<")
|
||||
if err != nil {
|
||||
msg, code := classifyApkOutput(string(out), err)
|
||||
return response{Error: msg, Code: code}
|
||||
|
||||
@@ -12,6 +12,13 @@ import (
|
||||
// Note: Command execution tests are not included here since apk is not available
|
||||
// in unit test environments. Integration tests would handle actual execution.
|
||||
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 {
|
||||
name string
|
||||
req request
|
||||
@@ -154,6 +161,13 @@ func TestValidPkgName(t *testing.T) {
|
||||
|
||||
// TestHandleRequest_AllActionsValidated tests both install and uninstall actions.
|
||||
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 {
|
||||
action string
|
||||
}{
|
||||
@@ -365,6 +379,12 @@ func TestValidPkgNameRegex_Compliance(t *testing.T) {
|
||||
|
||||
// TestHandleRequest_ErrorMessages tests that error messages are clear.
|
||||
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 {
|
||||
name string
|
||||
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),
|
||||
// but validation should pass.
|
||||
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 {
|
||||
action string
|
||||
pkg string
|
||||
@@ -438,6 +464,12 @@ func TestHandleRequest_SuccessPath(t *testing.T) {
|
||||
// TestHandleRequest_UpgradeValidation verifies that the upgrade action uses
|
||||
// the stricter validApkName regex (lowercase only, no @, no /).
|
||||
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 := []string{
|
||||
"curl",
|
||||
@@ -462,6 +494,12 @@ func TestHandleRequest_UpgradeValidation(t *testing.T) {
|
||||
|
||||
// TestHandleRequest_UpgradeInjectionPatterns verifies 5 injection patterns are rejected.
|
||||
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{
|
||||
"-malicious", // leading hyphen
|
||||
"pkg;evil", // semicolon
|
||||
@@ -486,6 +524,12 @@ func TestHandleRequest_UpgradeInjectionPatterns(t *testing.T) {
|
||||
// by legacy validPkgName for install/uninstall) is REJECTED by upgrade action
|
||||
// via the stricter validApkName.
|
||||
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{
|
||||
"pkg@edge", // @ accepted by validPkgName, rejected by validApkName
|
||||
"@scope/pkg", // npm scoped — rejected by validApkName
|
||||
|
||||
@@ -207,9 +207,9 @@ var coreToolSummaries = map[string]string{
|
||||
"session_status": "Show session status (model, tokens, compaction count)",
|
||||
"sessions_history": "Fetch message history for a 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_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",
|
||||
"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",
|
||||
|
||||
@@ -87,6 +87,7 @@ func (c *Config) ApplySystemConfigs(configs map[string]string) {
|
||||
str("tts.auto", &c.Tts.Auto)
|
||||
str("tts.mode", &c.Tts.Mode)
|
||||
integer("tts.max_length", &c.Tts.MaxLength)
|
||||
integer("tts.timeout_ms", &c.Tts.TimeoutMs)
|
||||
|
||||
// Cron
|
||||
integer("cron.max_retries", &c.Cron.MaxRetries)
|
||||
|
||||
@@ -222,10 +222,10 @@ func (cs *Service) GetJob(jobID string) (*Job, bool) {
|
||||
cs.mu.Lock()
|
||||
defer cs.mu.Unlock()
|
||||
|
||||
for i, job := range cs.store.Jobs {
|
||||
for _, job := range cs.store.Jobs {
|
||||
if job.ID == jobID {
|
||||
jobCopy := cs.store.Jobs[i]
|
||||
return &jobCopy, true
|
||||
result := job
|
||||
return &result, true
|
||||
}
|
||||
}
|
||||
return nil, false
|
||||
|
||||
@@ -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 ---
|
||||
|
||||
func TestService_AddJob_AtSchedule_DeleteAfterRun(t *testing.T) {
|
||||
@@ -235,18 +261,50 @@ func TestService_StartStop_JobExecution(t *testing.T) {
|
||||
|
||||
cs := NewService(storePath, handler)
|
||||
|
||||
interval := int64(50)
|
||||
_, err := cs.AddJob("fast", Schedule{Kind: "every", EveryMS: &interval}, "tick", false, "", "", "")
|
||||
if err := cs.Start(); err != nil {
|
||||
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 {
|
||||
t.Fatalf("AddJob error: %v", err)
|
||||
}
|
||||
|
||||
if err := cs.Start(); err != nil {
|
||||
t.Fatalf("Start error: %v", err)
|
||||
cs.mu.Lock()
|
||||
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
|
||||
time.Sleep(120 * time.Millisecond)
|
||||
deadline := time.Now().Add(500 * 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()
|
||||
|
||||
count := execCount.Load()
|
||||
|
||||
@@ -4,8 +4,11 @@ import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"fmt"
|
||||
"log/slog"
|
||||
"net"
|
||||
"regexp"
|
||||
"strconv"
|
||||
"strings"
|
||||
|
||||
"github.com/google/uuid"
|
||||
@@ -13,6 +16,7 @@ import (
|
||||
"github.com/nextlevelbuilder/goclaw/internal/gateway"
|
||||
"github.com/nextlevelbuilder/goclaw/internal/i18n"
|
||||
"github.com/nextlevelbuilder/goclaw/internal/permissions"
|
||||
"github.com/nextlevelbuilder/goclaw/internal/security"
|
||||
"github.com/nextlevelbuilder/goclaw/internal/store"
|
||||
"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
|
||||
// reference them directly.
|
||||
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.
|
||||
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
|
||||
// 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)")))
|
||||
return
|
||||
}
|
||||
if !bitrixDomainRegex.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")))
|
||||
if !bitrixCloudDomainRegex.MatchString(domain) && !selfHostedDomainRegex.MatchString(domain) {
|
||||
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
|
||||
}
|
||||
// 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 == "" {
|
||||
client.SendResponse(protocol.NewErrorResponse(req.ID, protocol.ErrInvalidRequest, i18n.T(locale, i18n.MsgRequired, "client_id and client_secret")))
|
||||
return
|
||||
@@ -352,6 +368,75 @@ func portalRowToView(row store.BitrixPortalData) bitrixPortalView {
|
||||
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
|
||||
// string substring match because the store interface doesn't expose typed
|
||||
// duplicate errors and we want consistent behaviour between pg + sqlite
|
||||
|
||||
@@ -4,6 +4,7 @@ import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"fmt"
|
||||
"strings"
|
||||
"sync"
|
||||
"testing"
|
||||
@@ -366,17 +367,62 @@ func TestBitrixPortals_Create_HappyPath_ReturnsInstallURL(t *testing.T) {
|
||||
func TestBitrixPortals_Create_InvalidDomain(t *testing.T) {
|
||||
tid := uuid.New()
|
||||
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{
|
||||
"name": "p",
|
||||
"domain": "not-a-bitrix-domain.com",
|
||||
"client_id": "x",
|
||||
"client_secret": "y",
|
||||
"name": "myportal",
|
||||
"domain": "myportal.bitrix24.com",
|
||||
"client_id": "local.abc",
|
||||
"client_secret": "secret123",
|
||||
}))
|
||||
|
||||
resp := readResponse(t, ch)
|
||||
if resp.Error == nil || resp.Error.Code != protocol.ErrInvalidRequest {
|
||||
t.Errorf("expected INVALID_REQUEST, got %+v", resp.Error)
|
||||
if resp.Error != nil {
|
||||
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",
|
||||
"a.bitrix24.com",
|
||||
"company.bitrix.info",
|
||||
"mycorp.bitrix24.vn",
|
||||
"portal.bitrix24.tr",
|
||||
"miempresa.bitrix24.es",
|
||||
"empresa.bitrix24.com.br",
|
||||
}
|
||||
bad := []string{
|
||||
"tamgiac.bitrix24",
|
||||
"tamgiac.example.com",
|
||||
"tamgiac.bitrix24.xx",
|
||||
"-bad.bitrix24.com",
|
||||
"UPPER.bitrix24.com", // we lowercase before match
|
||||
"a.b.bitrix24.com", // multi-level subdomain not allowed
|
||||
}
|
||||
for _, d := range good {
|
||||
if !bitrixDomainRegex.MatchString(d) {
|
||||
if !bitrixCloudDomainRegex.MatchString(d) {
|
||||
t.Errorf("should accept %q", d)
|
||||
}
|
||||
}
|
||||
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)
|
||||
}
|
||||
}
|
||||
@@ -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
|
||||
// to ALREADY_EXISTS.
|
||||
func TestIsDuplicateKeyErr(t *testing.T) {
|
||||
|
||||
@@ -15,8 +15,8 @@ import (
|
||||
|
||||
// ModelInfo is a normalized model entry returned by the list-models endpoint.
|
||||
type ModelInfo struct {
|
||||
ID string `json:"id"`
|
||||
Name string `json:"name,omitempty"`
|
||||
ID string `json:"id"`
|
||||
Name string `json:"name,omitempty"`
|
||||
Reasoning *providers.ReasoningCapability `json:"reasoning,omitempty"`
|
||||
}
|
||||
|
||||
@@ -109,11 +109,8 @@ func (h *ProvidersHandler) handleListProviderModels(w http.ResponseWriter, r *ht
|
||||
models = minimaxModels()
|
||||
default:
|
||||
// All other types use OpenAI-compatible /models endpoint
|
||||
apiBase := strings.TrimRight(h.resolveAPIBase(p), "/")
|
||||
if apiBase == "" {
|
||||
apiBase = "https://api.openai.com/v1"
|
||||
}
|
||||
models, err = fetchOpenAIModels(ctx, apiBase, p.APIKey)
|
||||
apiBase := openAIModelsAPIBase(p.ProviderType, h.resolveAPIBase(p))
|
||||
models, err = fetchOpenAIModels(ctx, apiBase, p.APIKey, openAIModelsExtraHeaders(p.ProviderType))
|
||||
}
|
||||
|
||||
if err != nil {
|
||||
@@ -126,6 +123,28 @@ func (h *ProvidersHandler) handleListProviderModels(w http.ResponseWriter, r *ht
|
||||
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(
|
||||
settings []byte,
|
||||
models []ModelInfo,
|
||||
|
||||
@@ -91,12 +91,15 @@ func fetchGeminiModels(ctx context.Context, apiKey string) ([]ModelInfo, error)
|
||||
}
|
||||
|
||||
// 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)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
req.Header.Set("Authorization", "Bearer "+apiKey)
|
||||
for k, v := range extraHeaders {
|
||||
req.Header.Set(k, v)
|
||||
}
|
||||
|
||||
resp, err := http.DefaultClient.Do(req)
|
||||
if err != nil {
|
||||
|
||||
@@ -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
|
||||
// handler fetches /api/tags from Ollama and maps rich details (family,
|
||||
// parameter_size, quantization_level) into the display name.
|
||||
|
||||
@@ -314,6 +314,18 @@ func (h *ProvidersHandler) registerInMemory(p *store.LLMProviderData) providerRu
|
||||
base = store.NovitaDefaultAPIBase
|
||||
}
|
||||
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:
|
||||
prov := providers.NewOpenAIProvider(p.Name, p.APIKey, apiBase, "")
|
||||
if p.ProviderType == store.ProviderMiniMax {
|
||||
|
||||
@@ -61,6 +61,10 @@ func (a *OpenAIAdapter) ToRequest(req ChatRequest) ([]byte, http.Header, error)
|
||||
if 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
|
||||
}
|
||||
|
||||
@@ -98,7 +98,7 @@ func (p *AnthropicProvider) buildRequestBody(model string, req ChatRequest, stre
|
||||
systemBlocks = append(systemBlocks, splitSystemPromptForCache(msg.Content)...)
|
||||
|
||||
case "user":
|
||||
if len(msg.Images) > 0 {
|
||||
if len(msg.Images) > 0 || len(msg.Videos) > 0 {
|
||||
var blocks []map[string]any
|
||||
for _, img := range msg.Images {
|
||||
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 != "" {
|
||||
blocks = append(blocks, map[string]any{
|
||||
"type": "text",
|
||||
|
||||
@@ -34,7 +34,7 @@ func (p *CodexProvider) buildRequestBody(req ChatRequest, stream bool) map[strin
|
||||
}
|
||||
|
||||
case "user":
|
||||
if len(m.Images) > 0 {
|
||||
if len(m.Images) > 0 || len(m.Videos) > 0 {
|
||||
var parts []map[string]any
|
||||
for _, img := range m.Images {
|
||||
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),
|
||||
})
|
||||
}
|
||||
// Videos are not supported by Codex, they are omitted here.
|
||||
if m.Content != "" {
|
||||
parts = append(parts, map[string]any{
|
||||
"type": "input_text",
|
||||
|
||||
@@ -17,6 +17,7 @@ type OpenAIProvider struct {
|
||||
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)
|
||||
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
|
||||
retryConfig RetryConfig
|
||||
middlewares RequestMiddleware // composed middleware chain (nil = no-op)
|
||||
@@ -63,6 +64,36 @@ func (p *OpenAIProvider) WithSiteInfo(url, title string) *OpenAIProvider {
|
||||
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.
|
||||
func (p *OpenAIProvider) WithRegistry(r ModelRegistry) *OpenAIProvider {
|
||||
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")
|
||||
}
|
||||
}
|
||||
@@ -45,6 +45,11 @@ func (p *OpenAIProvider) doRequest(ctx context.Context, body any) (io.ReadCloser
|
||||
if 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)
|
||||
if err != nil {
|
||||
|
||||
@@ -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.
|
||||
// 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
|
||||
// (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
|
||||
// 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 != "" {
|
||||
parts = append(parts, map[string]any{
|
||||
"type": "text",
|
||||
@@ -74,10 +86,26 @@ func (p *OpenAIProvider) buildRequestBody(model string, req ChatRequest, stream
|
||||
})
|
||||
}
|
||||
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{
|
||||
"type": "image_url",
|
||||
"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;
|
||||
// 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")
|
||||
// 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 {
|
||||
body["temperature"] = v
|
||||
}
|
||||
|
||||
@@ -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) {
|
||||
// Together behind reverse proxy — detected by providerType, not URL.
|
||||
p := NewOpenAIProvider("my-proxy", "key", "https://proxy.internal/v1", "")
|
||||
|
||||
@@ -113,13 +113,22 @@ type StreamChunk struct {
|
||||
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 {
|
||||
MimeType string `json:"mime_type"` // e.g. "image/jpeg"
|
||||
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)
|
||||
}
|
||||
|
||||
// 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.
|
||||
// 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).
|
||||
@@ -137,6 +146,7 @@ type Message struct {
|
||||
Content string `json:"content"`
|
||||
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)
|
||||
Videos []VideoContent `json:"-"` // vision: base64 videos (runtime only, never persisted to DB)
|
||||
MediaRefs []MediaRef `json:"media_refs,omitempty"` // persistent media file references
|
||||
ToolCalls []ToolCall `json:"tool_calls,omitempty"`
|
||||
ToolCallID string `json:"tool_call_id,omitempty"` // for role="tool" responses
|
||||
|
||||
@@ -6,6 +6,8 @@ import (
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"net"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"runtime"
|
||||
"strings"
|
||||
"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
|
||||
// parsed correctly into (ok, code, data, errMsg).
|
||||
func TestApkHelperCall_ValidResponse(t *testing.T) {
|
||||
|
||||
@@ -2,6 +2,7 @@ package skills
|
||||
|
||||
import (
|
||||
"bufio"
|
||||
"bytes"
|
||||
"context"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
@@ -39,6 +40,10 @@ const InstallTimeout = 5 * time.Minute
|
||||
// pkgHelperSocket is the Unix socket path for the root-privileged pkg-helper.
|
||||
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
|
||||
// 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
|
||||
@@ -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) {
|
||||
conn, err := net.DialTimeout("unix", pkgHelperSocket, 5*time.Second)
|
||||
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)
|
||||
}
|
||||
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"
|
||||
}
|
||||
|
||||
var resp struct {
|
||||
OK bool `json:"ok"`
|
||||
Error string `json:"error"`
|
||||
Code string `json:"code"`
|
||||
Data string `json:"data"`
|
||||
}
|
||||
if err := json.Unmarshal(scanner.Bytes(), &resp); err != nil {
|
||||
resp, err := parsePkgHelperResponse(scanner.Bytes())
|
||||
if err != nil {
|
||||
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.
|
||||
if resp.Code == "" && !resp.OK {
|
||||
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,
|
||||
|
||||
@@ -34,6 +34,7 @@ const (
|
||||
ProviderBytePlus = "byteplus" // BytePlus ModelArk (Seed 2.0 models)
|
||||
ProviderBytePlusCoding = "byteplus_coding" // BytePlus ModelArk Coding Plan
|
||||
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.
|
||||
NovitaDefaultAPIBase = "https://api.novita.ai/openai"
|
||||
@@ -44,6 +45,12 @@ const (
|
||||
BytePlusCodingDefaultAPIBase = "https://ark.ap-southeast.bytepluses.com/api/coding/v3"
|
||||
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
|
||||
@@ -77,6 +84,7 @@ var ValidProviderTypes = map[string]bool{
|
||||
ProviderBytePlus: true,
|
||||
ProviderBytePlusCoding: true,
|
||||
ProviderVertex: true,
|
||||
ProviderKimiCoding: true,
|
||||
}
|
||||
|
||||
// VertexProviderSettings holds Vertex-specific config stored in llm_providers.settings JSONB.
|
||||
|
||||
@@ -264,3 +264,153 @@ func parseGeminiResponse(respBody []byte) (*providers.ChatResponse, error) {
|
||||
},
|
||||
}, 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)
|
||||
}
|
||||
@@ -10,6 +10,7 @@ import (
|
||||
"strings"
|
||||
|
||||
"github.com/nextlevelbuilder/goclaw/internal/providers"
|
||||
"github.com/nextlevelbuilder/goclaw/internal/security"
|
||||
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) 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 {
|
||||
@@ -77,6 +78,10 @@ func (t *ReadImageTool) Parameters() map[string]any {
|
||||
"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.",
|
||||
},
|
||||
"url": map[string]any{
|
||||
"type": "string",
|
||||
"description": "Optional URL to an image. Use this to analyze images hosted online.",
|
||||
},
|
||||
},
|
||||
"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."
|
||||
}
|
||||
|
||||
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
|
||||
images := MediaImagesFromCtx(ctx)
|
||||
if imgPath, _ := args["path"].(string); imgPath != "" {
|
||||
if imgPath != "" {
|
||||
fileImages, err := t.loadImageFromPath(ctx, imgPath)
|
||||
if err != nil {
|
||||
return ErrorResult(err.Error())
|
||||
}
|
||||
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 {
|
||||
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", "", "",
|
||||
@@ -138,6 +157,24 @@ func (t *ReadImageTool) callProvider(ctx context.Context, cp credentialProvider,
|
||||
prompt := GetParamString(params, "prompt", "Describe this image in detail.")
|
||||
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
|
||||
p, err := t.registry.Get(ctx, providerName)
|
||||
if err != nil {
|
||||
|
||||
@@ -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")
|
||||
}
|
||||
}
|
||||
@@ -4,10 +4,13 @@ import (
|
||||
"context"
|
||||
"fmt"
|
||||
"log/slog"
|
||||
"net"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"strings"
|
||||
|
||||
"github.com/nextlevelbuilder/goclaw/internal/providers"
|
||||
"github.com/nextlevelbuilder/goclaw/internal/security"
|
||||
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).
|
||||
const videoMaxBytes = 100 * 1024 * 1024
|
||||
|
||||
const videoURLPinnedIPParam = "_pinned_ip"
|
||||
|
||||
// videoProviderPriority is the order in which providers are tried for video analysis.
|
||||
// OpenAI excluded — no native video upload in chat completions.
|
||||
var videoProviderPriority = []string{"gemini", "openrouter"}
|
||||
@@ -77,6 +82,10 @@ func (t *ReadVideoTool) Parameters() map[string]any {
|
||||
"type": "string",
|
||||
"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"},
|
||||
}
|
||||
@@ -88,21 +97,43 @@ func (t *ReadVideoTool) Execute(ctx context.Context, args map[string]any) *Resul
|
||||
prompt = "Analyze this video and describe its contents."
|
||||
}
|
||||
mediaID, _ := args["media_id"].(string)
|
||||
videoURL, _ := args["url"].(string)
|
||||
|
||||
videoPath, videoMime, err := t.resolveVideoFile(ctx, mediaID)
|
||||
if err != nil {
|
||||
return ErrorResult(err.Error())
|
||||
if mediaID != "" && videoURL != "" {
|
||||
return ErrorResult("Both 'media_id' and 'url' parameters cannot be specified. Choose only one.")
|
||||
}
|
||||
|
||||
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 err != nil {
|
||||
return ErrorResult(fmt.Sprintf("Failed to read video file: %v", err))
|
||||
}
|
||||
slog.Info("read_video: file loaded", "size_bytes", len(data))
|
||||
if len(data) > videoMaxBytes {
|
||||
return ErrorResult(fmt.Sprintf("Video too large: %d bytes (max %d)", len(data), videoMaxBytes))
|
||||
if videoURL != "" {
|
||||
validatedURL, validatedIP, err := security.Validate(videoURL)
|
||||
if err != nil {
|
||||
return ErrorResult(fmt.Sprintf("Invalid video URL: %v", err))
|
||||
}
|
||||
pinnedIP = validatedIP
|
||||
|
||||
// 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", "", "",
|
||||
@@ -114,7 +145,11 @@ func (t *ReadVideoTool) Execute(ctx context.Context, args map[string]any) *Resul
|
||||
}
|
||||
chain[i].Params["prompt"] = prompt
|
||||
chain[i].Params["data"] = data
|
||||
chain[i].Params["url"] = videoURL
|
||||
chain[i].Params["mime"] = videoMime
|
||||
if pinnedIP != nil {
|
||||
chain[i].Params[videoURLPinnedIPParam] = pinnedIP
|
||||
}
|
||||
}
|
||||
|
||||
chainResult, err := ExecuteWithChain(ctx, chain, t.registry, t.callProvider)
|
||||
|
||||
@@ -5,10 +5,14 @@ import (
|
||||
"encoding/base64"
|
||||
"fmt"
|
||||
"log/slog"
|
||||
"net"
|
||||
"net/http"
|
||||
"path/filepath"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"github.com/nextlevelbuilder/goclaw/internal/providers"
|
||||
"github.com/nextlevelbuilder/goclaw/internal/security"
|
||||
)
|
||||
|
||||
// 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.
|
||||
// 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) {
|
||||
prompt := GetParamString(params, "prompt", "Analyze this video and describe its contents.")
|
||||
data, _ := params["data"].([]byte)
|
||||
videoURL, _ := params["url"].(string)
|
||||
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).
|
||||
ptype := GetParamString(params, "_provider_type", providerTypeFromName(providerName))
|
||||
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{
|
||||
Messages: []providers.Message{{Role: "user", Content: prompt}},
|
||||
Model: model,
|
||||
@@ -79,7 +93,71 @@ func (t *ReadVideoTool) callProvider(ctx context.Context, cp credentialProvider,
|
||||
if reserveErr != nil {
|
||||
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 {
|
||||
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
|
||||
}
|
||||
|
||||
// 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)
|
||||
if err != nil {
|
||||
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{
|
||||
Messages: []providers.Message{
|
||||
{
|
||||
Role: "user",
|
||||
Content: prompt,
|
||||
Images: []providers.ImageContent{{MimeType: mime, Data: base64.StdEncoding.EncodeToString(data)}},
|
||||
Videos: []providers.VideoContent{vidContent},
|
||||
},
|
||||
},
|
||||
Model: model,
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
@@ -254,7 +254,7 @@ func resolveRemoteCDP(remoteURL string) (string, error) {
|
||||
defer resp.Body.Close()
|
||||
|
||||
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 {
|
||||
|
||||
@@ -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: "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: "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_cloud", label: "Ollama Cloud", apiBase: "https://ollama.com/v1", placeholder: "" },
|
||||
{ value: "claude_cli", label: "Claude CLI (Local)", apiBase: "", placeholder: "" },
|
||||
|
||||
@@ -721,7 +721,7 @@
|
||||
},
|
||||
"errors": {
|
||||
"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.",
|
||||
"forbidden": "You need tenant admin permission to create portals.",
|
||||
"gatewayURLUnknown": "Open the goclaw UI via your public URL first (not localhost), then retry."
|
||||
|
||||
@@ -639,7 +639,7 @@
|
||||
},
|
||||
"errors": {
|
||||
"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.",
|
||||
"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."
|
||||
|
||||
@@ -639,7 +639,7 @@
|
||||
},
|
||||
"errors": {
|
||||
"invalidName": "请使用小写字母、数字、连字符、下划线(2-64 字符)。",
|
||||
"invalidDomain": "必须是有效的 Bitrix24 域名(例如:mycorp.bitrix24.com)。",
|
||||
"invalidDomain": "必须是有效的主机名(例如:mycorp.bitrix24.com、mycorp.bitrix24.vn 或你的自托管域名)。",
|
||||
"duplicateName": "门户名称已存在。",
|
||||
"forbidden": "您需要租户管理员权限才能创建门户。",
|
||||
"gatewayURLUnknown": "请先通过公网 URL(非 localhost)打开 goclaw UI 然后重试。"
|
||||
|
||||
@@ -10,13 +10,72 @@ import { useBitrixPortalCreate } from "./use-bitrix-portals";
|
||||
// Validation mirrors the server-side regex in
|
||||
// internal/gateway/methods/bitrix_portals.go. Server is authoritative;
|
||||
// client validation is purely UX so the operator gets feedback before a
|
||||
// round-trip. Pattern intentionally accepts a wide TLD set (Bitrix24 has
|
||||
// regional clouds: .com, .eu, .ru, .de, .fr, .jp, .in, .kz, .ua, .by) plus
|
||||
// .bitrix.info for self-hosted.
|
||||
// round-trip. Pattern accepts Bitrix24 regional clouds (.com, .eu, .ru,
|
||||
// .de, .fr, .jp, .in, .kz, .ua, .by, .vn, .tr, .es, .com.br, .com.ar),
|
||||
// .bitrix.info self-hosted, plus any valid hostname/port for fully
|
||||
// self-hosted custom domains.
|
||||
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]$/;
|
||||
|
||||
// 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 {
|
||||
/** Invoked with the server response after bitrix.portals.create succeeds. */
|
||||
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).",
|
||||
});
|
||||
}
|
||||
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", {
|
||||
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 (!clientSecret.trim()) e.client_secret = t("common.required", { defaultValue: "Required" });
|
||||
@@ -121,7 +191,7 @@ export function BitrixPortalFormStep({ onSuccess, onCancel }: BitrixPortalFormSt
|
||||
value={domain}
|
||||
onChange={(e) => setDomain(e.target.value)}
|
||||
onBlur={handleDomainBlur}
|
||||
placeholder="tamgiac.bitrix24.com"
|
||||
placeholder="mycorp.bitrix24.vn or bitrix.example.com"
|
||||
autoComplete="off"
|
||||
autoFocus
|
||||
/>
|
||||
|
||||
@@ -4,8 +4,8 @@ export function SetupLayout({ children }: { children: React.ReactNode }) {
|
||||
const { t } = useTranslation("setup");
|
||||
|
||||
return (
|
||||
<div className="flex min-h-dvh items-center justify-center bg-background px-4 py-8">
|
||||
<div className="w-full max-w-2xl space-y-6">
|
||||
<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 overflow-y-auto max-h-dvh sm:max-h-none">
|
||||
<div className="text-center">
|
||||
<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>
|
||||
|
||||
Reference in new issue
Block a user