diff --git a/internal/mcp/server_test.go b/internal/mcp/server_test.go index 4ed0d47e..8ec5e8b3 100644 --- a/internal/mcp/server_test.go +++ b/internal/mcp/server_test.go @@ -6,6 +6,7 @@ import ( "net/http" "net/http/httptest" "path/filepath" + "slices" "strings" "testing" @@ -98,8 +99,21 @@ func TestOpenAIFileMetadataMatchesDeclaredSchemas(t *testing.T) { if !ok { t.Fatalf("%s file input path %q missing from input schema", def.Name, path) } - if property["type"] != "object" || property["additionalProperties"] != true { - t.Fatalf("%s file input %q must be an open object: %#v", def.Name, path, property) + if property["type"] != "object" { + t.Fatalf("%s file input %q must be an object: %#v", def.Name, path, property) + } + properties, ok := property["properties"].(map[string]any) + if !ok { + t.Fatalf("%s file input %q properties = %#v", def.Name, path, property["properties"]) + } + for _, key := range []string{"download_url", "file_id", "file_name", "mime_type"} { + if _, ok := properties[key]; !ok { + t.Fatalf("%s file input %q missing documented property %q", def.Name, path, key) + } + } + required, ok := property["required"].([]string) + if !ok || !slices.Contains(required, "download_url") || !slices.Contains(required, "file_id") { + t.Fatalf("%s file input %q required = %#v", def.Name, path, property["required"]) } } if len(def.FileArgRewritePaths) > 0 { diff --git a/internal/tool/media/contract.go b/internal/tool/media/contract.go index 9507c6dd..cbccdd70 100644 --- a/internal/tool/media/contract.go +++ b/internal/tool/media/contract.go @@ -52,7 +52,18 @@ func InputSchema(name string) (map[string]any, bool) { } return schema, true case ToolFilePublish: - props["file"] = toolcontract.OpenObject("Top-level connector file object. Connector runtimes may rewrite this object with mounted local-path metadata before the tool call reaches AgentDock.") + props["file"] = map[string]any{ + "type": "object", + "description": "OpenAI connector file object. download_url and file_id are required; file_name and mime_type are optional.", + "additionalProperties": true, + "properties": map[string]any{ + "download_url": stringProp("Temporary URL used to download the connector-provided file."), + "file_id": stringProp("Stable connector file identifier."), + "file_name": stringProp("Optional original file name."), + "mime_type": stringProp("Optional media type for the file."), + }, + "required": []string{"download_url", "file_id"}, + } props["path"] = stringProp("Local file or directory path visible to this AgentDock instance. Relative paths resolve from ~/AgentDock.") props["retention_seconds"] = boundedIntProp("Signed URL retention seconds. Zero uses the default 86400 and values are capped at 604800.", 0, int(publicartifacts.MaxRetention/time.Second)) default: diff --git a/internal/tool/media/file_publish.go b/internal/tool/media/file_publish.go index 5fc16910..49293ba5 100644 --- a/internal/tool/media/file_publish.go +++ b/internal/tool/media/file_publish.go @@ -3,18 +3,35 @@ package media import ( "context" "fmt" + "io" + "net" + "net/http" + "net/url" + "os" + "path" "path/filepath" "strings" + "time" "github.com/uvwt/agentdock/internal/httpx/requestmeta" "github.com/uvwt/agentdock/internal/publicartifacts" ) +const maxConnectorFileBytes int64 = 128 << 20 + +type connectorFileInput struct { + DownloadURL string + FileID string + FileName string + MimeType string +} + func (s *Service) FilePublish(ctx context.Context, request FilePublishRequest) (Result, error) { - pathValue, err := s.filePublishSourcePath(request) + pathValue, cleanup, err := s.filePublishSourcePath(ctx, request) if err != nil { return nil, err } + defer cleanup() store := publicartifacts.New(s.cfg.AgentDockHome, s.cfg.OAuthServerURL, s.cfg.Port) published, err := store.Publish(publicartifacts.PublishRequest{Path: pathValue, RetentionSeconds: intValue(request.RetentionSeconds, 0), BaseURL: requestmeta.BaseURL(ctx)}) if err != nil { @@ -27,25 +44,28 @@ func (s *Service) FilePublish(ctx context.Context, request FilePublishRequest) ( return result, nil } -func (s *Service) filePublishSourcePath(request FilePublishRequest) (string, error) { +func (s *Service) filePublishSourcePath(ctx context.Context, request FilePublishRequest) (string, func(), error) { if request.File != nil { if pathValue := connectorLocalPath(request.File); pathValue != "" { resolved, err := s.ws.ResolveExisting(pathValue) if err != nil { - return "", err + return "", func() {}, err } - return resolved.Abs, nil + return resolved.Abs, func() {}, nil + } + if input, ok := connectorFile(request.File); ok { + return s.downloadConnectorFile(ctx, input, newConnectorDownloadClient()) } } pathValue := strings.TrimSpace(request.Path) if pathValue == "" { - return "", toolError("FILE_PUBLISH_SOURCE_REQUIRED", "file or path is required", "validation") + return "", func() {}, toolError("FILE_PUBLISH_SOURCE_REQUIRED", "file or path is required", "validation") } resolved, err := s.ws.ResolveExisting(pathValue) if err != nil { - return "", err + return "", func() {}, err } - return resolved.Abs, nil + return resolved.Abs, func() {}, nil } func connectorLocalPath(value any) string { @@ -64,3 +84,186 @@ func connectorLocalPath(value any) string { } return "" } + +func connectorFile(value any) (connectorFileInput, bool) { + v, ok := value.(map[string]any) + if !ok { + return connectorFileInput{}, false + } + input := connectorFileInput{ + DownloadURL: stringValue(v["download_url"]), + FileID: stringValue(v["file_id"]), + FileName: stringValue(v["file_name"]), + MimeType: stringValue(v["mime_type"]), + } + if input.DownloadURL == "" && input.FileID == "" { + return connectorFileInput{}, false + } + return input, true +} + +func stringValue(value any) string { + text, _ := value.(string) + return strings.TrimSpace(text) +} + +func (s *Service) downloadConnectorFile(ctx context.Context, input connectorFileInput, client *http.Client) (string, func(), error) { + if input.DownloadURL == "" || input.FileID == "" { + return "", func() {}, toolError("FILE_PUBLISH_FILE_INVALID", "file.download_url and file.file_id are required", "validation") + } + parsed, err := validateConnectorDownloadURL(input.DownloadURL) + if err != nil { + return "", func() {}, err + } + tmpRoot := filepath.Join(s.cfg.AgentDockHome, "tmp") + if err := os.MkdirAll(tmpRoot, 0o700); err != nil { + return "", func() {}, fmt.Errorf("create file publish temp root: %w", err) + } + tmpDir, err := os.MkdirTemp(tmpRoot, "file-publish-") + if err != nil { + return "", func() {}, fmt.Errorf("create file publish temp dir: %w", err) + } + cleanup := func() { _ = os.RemoveAll(tmpDir) } + filename := connectorDownloadName(input, parsed) + target := filepath.Join(tmpDir, filename) + + req, err := http.NewRequestWithContext(ctx, http.MethodGet, input.DownloadURL, nil) + if err != nil { + cleanup() + return "", func() {}, toolError("FILE_PUBLISH_DOWNLOAD_URL_INVALID", "cannot create file download request", "validation") + } + resp, err := client.Do(req) + if err != nil { + cleanup() + return "", func() {}, toolError("FILE_PUBLISH_DOWNLOAD_FAILED", "cannot download connector file", "runtime") + } + defer resp.Body.Close() + if resp.StatusCode < 200 || resp.StatusCode >= 300 { + cleanup() + return "", func() {}, toolError("FILE_PUBLISH_DOWNLOAD_HTTP_ERROR", fmt.Sprintf("connector file download returned HTTP %d", resp.StatusCode), "runtime") + } + if resp.ContentLength > maxConnectorFileBytes { + cleanup() + return "", func() {}, toolError("FILE_PUBLISH_FILE_TOO_LARGE", "connector file exceeds maximum supported size", "validation") + } + + f, err := os.OpenFile(target, os.O_CREATE|os.O_EXCL|os.O_WRONLY, 0o600) + if err != nil { + cleanup() + return "", func() {}, fmt.Errorf("create connector temp file: %w", err) + } + written, copyErr := io.Copy(f, io.LimitReader(resp.Body, maxConnectorFileBytes+1)) + closeErr := f.Close() + if copyErr != nil { + cleanup() + return "", func() {}, toolError("FILE_PUBLISH_DOWNLOAD_FAILED", "cannot read connector file download", "runtime") + } + if closeErr != nil { + cleanup() + return "", func() {}, fmt.Errorf("close connector temp file: %w", closeErr) + } + if written > maxConnectorFileBytes { + cleanup() + return "", func() {}, toolError("FILE_PUBLISH_FILE_TOO_LARGE", "connector file exceeds maximum supported size", "validation") + } + return target, cleanup, nil +} + +func validateConnectorDownloadURL(rawURL string) (*url.URL, error) { + parsed, err := url.Parse(rawURL) + if err != nil || parsed.Host == "" { + return nil, toolError("FILE_PUBLISH_DOWNLOAD_URL_INVALID", "file.download_url must be an absolute HTTP(S) URL", "validation") + } + parsed.Scheme = strings.ToLower(parsed.Scheme) + if parsed.Scheme != "https" && parsed.Scheme != "http" { + return nil, toolError("FILE_PUBLISH_DOWNLOAD_URL_INVALID", "file.download_url must use http or https", "validation") + } + return parsed, nil +} + +func connectorDownloadName(input connectorFileInput, parsed *url.URL) string { + name := strings.TrimSpace(input.FileName) + if name == "" { + name = path.Base(parsed.Path) + } + name = strings.ReplaceAll(name, "\\", "/") + name = path.Base(name) + if name == "" || name == "." || name == "/" { + name = "uploaded-file" + } + var b strings.Builder + for _, r := range name { + if r < 0x20 || r == 0x7f { + b.WriteByte('_') + continue + } + b.WriteRune(r) + } + name = strings.TrimSpace(b.String()) + if name == "" || name == "." || name == ".." { + return "uploaded-file" + } + return name +} + +func newConnectorDownloadClient() *http.Client { + transport := http.DefaultTransport.(*http.Transport).Clone() + transport.Proxy = nil + transport.DialContext = safeConnectorDialContext + return &http.Client{ + Transport: transport, + Timeout: 2 * time.Minute, + CheckRedirect: func(req *http.Request, via []*http.Request) error { + if len(via) >= 5 { + return fmt.Errorf("too many redirects") + } + scheme := strings.ToLower(req.URL.Scheme) + if scheme != "http" && scheme != "https" { + return fmt.Errorf("redirect uses unsupported scheme %q", req.URL.Scheme) + } + return nil + }, + } +} + +func safeConnectorDialContext(ctx context.Context, network, address string) (net.Conn, error) { + host, port, err := net.SplitHostPort(address) + if err != nil { + return nil, err + } + addresses, err := net.DefaultResolver.LookupIPAddr(ctx, host) + if err != nil { + return nil, err + } + if len(addresses) == 0 { + return nil, fmt.Errorf("connector download host resolved to no addresses") + } + for _, resolved := range addresses { + if blockedConnectorDownloadIP(resolved.IP) { + return nil, fmt.Errorf("connector download host resolves to a non-public address") + } + } + dialer := &net.Dialer{Timeout: 15 * time.Second, KeepAlive: 30 * time.Second} + var lastErr error + for _, resolved := range addresses { + conn, dialErr := dialer.DialContext(ctx, network, net.JoinHostPort(resolved.IP.String(), port)) + if dialErr == nil { + return conn, nil + } + lastErr = dialErr + } + return nil, lastErr +} + +func blockedConnectorDownloadIP(ip net.IP) bool { + if ip == nil { + return true + } + if ip.IsLoopback() || ip.IsPrivate() || ip.IsUnspecified() || ip.IsMulticast() || ip.IsLinkLocalUnicast() || ip.IsLinkLocalMulticast() { + return true + } + if ip4 := ip.To4(); ip4 != nil { + return ip4[0] == 100 && ip4[1]&0xc0 == 0x40 + } + return false +} diff --git a/internal/tool/media/file_publish_test.go b/internal/tool/media/file_publish_test.go new file mode 100644 index 00000000..f98e4167 --- /dev/null +++ b/internal/tool/media/file_publish_test.go @@ -0,0 +1,107 @@ +package media + +import ( + "context" + "net" + "net/http" + "net/http/httptest" + "os" + "path/filepath" + "slices" + "testing" + + "github.com/uvwt/agentdock/internal/config" + "github.com/uvwt/agentdock/internal/workspace" +) + +func TestFilePublishInputSchemaUsesDocumentedOpenAIFileShape(t *testing.T) { + schema, ok := InputSchema(ToolFilePublish) + if !ok { + t.Fatal("file_publish input schema missing") + } + properties := schema["properties"].(map[string]any) + file := properties["file"].(map[string]any) + if file["type"] != "object" { + t.Fatalf("file.type = %#v, want object", file["type"]) + } + fileProperties := file["properties"].(map[string]any) + for _, key := range []string{"download_url", "file_id", "file_name", "mime_type"} { + if _, ok := fileProperties[key]; !ok { + t.Fatalf("file.properties missing %q", key) + } + } + required, ok := file["required"].([]string) + if !ok { + t.Fatalf("file.required = %#v", file["required"]) + } + for _, key := range []string{"download_url", "file_id"} { + if !slices.Contains(required, key) { + t.Fatalf("file.required = %#v, missing %q", required, key) + } + } +} + +func TestDownloadConnectorFileUsesDocumentedOpenAIFileObject(t *testing.T) { + payload := []byte("connector file payload") + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + w.Header().Set("Content-Type", "text/plain") + _, _ = w.Write(payload) + })) + defer server.Close() + + home := t.TempDir() + root := filepath.Join(t.TempDir(), "workspace") + ws, err := workspace.New(root) + if err != nil { + t.Fatal(err) + } + service := New(config.Config{AgentDockHome: home}, ws, nil) + input := connectorFileInput{ + DownloadURL: server.URL + "/download", + FileID: "file-test-123", + FileName: "../review.txt", + MimeType: "text/plain", + } + pathValue, cleanup, err := service.downloadConnectorFile(context.Background(), input, server.Client()) + if err != nil { + t.Fatal(err) + } + if filepath.Base(pathValue) != "review.txt" { + t.Fatalf("downloaded basename = %q", filepath.Base(pathValue)) + } + got, err := os.ReadFile(pathValue) + if err != nil { + t.Fatal(err) + } + if string(got) != string(payload) { + t.Fatalf("downloaded payload = %q", got) + } + cleanup() + if _, err := os.Stat(pathValue); !os.IsNotExist(err) { + t.Fatalf("temp file still exists after cleanup: %v", err) + } +} + +func TestBlockedConnectorDownloadIP(t *testing.T) { + blocked := []string{ + "127.0.0.1", + "10.0.0.1", + "172.16.0.1", + "192.168.1.1", + "169.254.169.254", + "100.64.0.1", + "::1", + "fc00::1", + "fe80::1", + } + for _, raw := range blocked { + if !blockedConnectorDownloadIP(net.ParseIP(raw)) { + t.Errorf("blockedConnectorDownloadIP(%q) = false", raw) + } + } + for _, raw := range []string{"8.8.8.8", "1.1.1.1", "2606:4700:4700::1111"} { + if blockedConnectorDownloadIP(net.ParseIP(raw)) { + t.Errorf("blockedConnectorDownloadIP(%q) = true", raw) + } + } +}