From 18e698bed6bf98e9affda25b3a82314e31e489ed Mon Sep 17 00:00:00 2001 From: mt21625457 Date: Fri, 17 Jul 2026 11:45:14 +0800 Subject: [PATCH] feat(security-audit): allow admin-managed audit node targets and polish pool UI Let admins configure private/intranet Guard endpoints without destination-class blocking, and fix prompt-audit switch layout so thumbs and labels no longer overlap. Co-authored-by: Cursor --- .../securityaudit/prompt_outbound_security.go | 139 +---------------- .../prompt_outbound_security_test.go | 52 +++---- deploy/build_image.sh | 0 deploy/docker-compose.yml | 2 +- .../features/prompt-audit/PromptAuditView.vue | 142 +++++++++++++----- .../__tests__/PromptAuditView.spec.ts | 30 ++++ .../prompt-audit/components/EndpointPool.vue | 135 ++++++++++------- .../components/EventWorkspace.vue | 18 +-- .../prompt-audit/components/PolicyPanel.vue | 12 +- .../components/RuntimeOverview.vue | 100 ++++++------ .../src/i18n/locales/en/admin/promptAudit.ts | 2 + .../src/i18n/locales/zh/admin/promptAudit.ts | 2 + frontend/src/style.css | 8 + .../design.md | 14 +- .../specs/prompt-input-audit/spec.md | 18 ++- .../tasks.md | 13 ++ 16 files changed, 342 insertions(+), 345 deletions(-) mode change 100644 => 100755 deploy/build_image.sh diff --git a/backend/internal/securityaudit/prompt_outbound_security.go b/backend/internal/securityaudit/prompt_outbound_security.go index f1e3ac3fc..81e987eb4 100644 --- a/backend/internal/securityaudit/prompt_outbound_security.go +++ b/backend/internal/securityaudit/prompt_outbound_security.go @@ -1,13 +1,9 @@ package securityaudit import ( - "context" "crypto/tls" - "errors" - "fmt" "net" "net/http" - "net/netip" "net/url" "strings" "time" @@ -17,40 +13,6 @@ import ( const maxGuardResponseBytes int64 = 256 * 1024 -var ( - errRedirectBlocked = errors.New("prompt guard redirect blocked") - metadataHosts = map[string]struct{}{ - "metadata": {}, "metadata.google.internal": {}, "metadata.azure.internal": {}, - "instance-data": {}, "instance-data.ec2.internal": {}, - } - blockedPrefixes = []netip.Prefix{ - netip.MustParsePrefix("0.0.0.0/8"), - netip.MustParsePrefix("100.64.0.0/10"), - netip.MustParsePrefix("169.254.0.0/16"), - netip.MustParsePrefix("192.0.0.0/24"), - netip.MustParsePrefix("192.0.2.0/24"), - netip.MustParsePrefix("198.18.0.0/15"), - netip.MustParsePrefix("198.51.100.0/24"), - netip.MustParsePrefix("203.0.113.0/24"), - netip.MustParsePrefix("224.0.0.0/4"), - netip.MustParsePrefix("240.0.0.0/4"), - netip.MustParsePrefix("::/128"), - netip.MustParsePrefix("fe80::/10"), - netip.MustParsePrefix("ff00::/8"), - netip.MustParsePrefix("2001:db8::/32"), - } -) - -type DNSResolver interface { - LookupNetIP(ctx context.Context, network, host string) ([]netip.Addr, error) -} - -type netResolver struct{ resolver *net.Resolver } - -func (r netResolver) LookupNetIP(ctx context.Context, network, host string) ([]netip.Addr, error) { - return r.resolver.LookupNetIP(ctx, network, host) -} - func NormalizeBaseURL(raw string) (string, error) { raw = strings.TrimSpace(raw) parsed, err := url.Parse(raw) @@ -64,29 +26,10 @@ func NormalizeBaseURL(raw string) (string, error) { if parsed.User != nil || parsed.RawQuery != "" || parsed.Fragment != "" { return "", infraerrors.BadRequest("prompt_audit_unsafe_base_url", "审计节点地址不能包含凭据、查询参数或片段") } - host := strings.ToLower(strings.TrimSuffix(parsed.Hostname(), ".")) + host := strings.TrimSpace(parsed.Hostname()) if host == "" { return "", infraerrors.BadRequest("prompt_audit_invalid_base_url", "审计节点地址无效") } - if _, blocked := metadataHosts[host]; blocked || strings.HasSuffix(host, ".metadata.google.internal") { - return "", infraerrors.BadRequest("prompt_audit_unsafe_base_url", "审计节点地址不在允许范围") - } - allowPrivate := isExplicitPrivateHost(host) - if addr, err := netip.ParseAddr(host); err == nil { - if isBlockedAddress(addr) { - return "", infraerrors.BadRequest("prompt_audit_unsafe_base_url", "审计节点地址不在允许范围") - } - // Loopback literals remain available for local Guard nodes and tests. - // RFC1918 literals are rejected so an admin session cannot pivot into - // arbitrary private-network services; use a hostname allowlist instead. - if addr.IsPrivate() { - return "", infraerrors.BadRequest("prompt_audit_unsafe_base_url", "审计节点地址不在允许范围") - } - allowPrivate = addr.IsLoopback() - } - if parsed.Scheme == "http" && !allowPrivate { - return "", infraerrors.BadRequest("prompt_audit_https_required", "公网审计节点必须使用 HTTPS") - } path := strings.TrimRight(parsed.EscapedPath(), "/") if strings.EqualFold(path, "/v1") { path = "" @@ -113,17 +56,10 @@ func ModelsURL(base string) (string, error) { } func NewSecureHTTPClient(endpoint ActiveEndpoint) (*http.Client, error) { - normalized, err := NormalizeBaseURL(endpoint.BaseURL) + _, err := NormalizeBaseURL(endpoint.BaseURL) if err != nil { return nil, err } - parsed, _ := url.Parse(normalized) - host := strings.ToLower(strings.TrimSuffix(parsed.Hostname(), ".")) - allowPrivate := isExplicitPrivateHost(host) - if addr, parseErr := netip.ParseAddr(host); parseErr == nil { - allowPrivate = addr.IsLoopback() - } - resolver := netResolver{resolver: net.DefaultResolver} dialer := &net.Dialer{Timeout: 3 * time.Second, KeepAlive: 30 * time.Second} transport := &http.Transport{ // Do not inherit HTTP(S)_PROXY. A proxy would move the actual destination @@ -138,7 +74,10 @@ func NewSecureHTTPClient(endpoint ActiveEndpoint) (*http.Client, error) { ExpectContinueTimeout: time.Second, TLSClientConfig: &tls.Config{MinVersion: tls.VersionTLS12}, } - transport.DialContext = secureDialContext(dialer, resolver, allowPrivate) + // Endpoint ownership and destination trust are administrator concerns. + // Use the standard dialer so configured private, loopback, reserved, and + // DNS-resolved addresses are all reachable from the service environment. + transport.DialContext = dialer.DialContext timeout := time.Duration(endpoint.TimeoutMS) * time.Millisecond if timeout <= 0 { timeout = DefaultTimeoutMS * time.Millisecond @@ -146,71 +85,5 @@ func NewSecureHTTPClient(endpoint ActiveEndpoint) (*http.Client, error) { return &http.Client{ Transport: transport, Timeout: timeout, - CheckRedirect: func(_ *http.Request, _ []*http.Request) error { - return errRedirectBlocked - }, }, nil } - -func secureDialContext(dialer *net.Dialer, resolver DNSResolver, allowPrivate bool) func(context.Context, string, string) (net.Conn, error) { - return func(ctx context.Context, network, address string) (net.Conn, error) { - host, port, err := net.SplitHostPort(address) - if err != nil { - return nil, fmt.Errorf("prompt guard dial address invalid") - } - addresses, err := resolver.LookupNetIP(ctx, "ip", host) - if err != nil || len(addresses) == 0 { - return nil, fmt.Errorf("prompt guard dns unavailable") - } - var lastErr error - for _, addr := range addresses { - if isBlockedAddress(addr) { - lastErr = fmt.Errorf("prompt guard resolved address blocked") - continue - } - if allowPrivate { - // localhost / *.localhost may only resolve to loopback. A hosts or - // DNS mapping from localhost to RFC1918 must not become an SSRF pivot. - if !addr.IsLoopback() { - lastErr = fmt.Errorf("prompt guard resolved address blocked") - continue - } - } else if addr.IsPrivate() || addr.IsLoopback() { - lastErr = fmt.Errorf("prompt guard resolved address blocked") - continue - } - if !addr.IsGlobalUnicast() && !addr.IsLoopback() { - lastErr = fmt.Errorf("prompt guard resolved address blocked") - continue - } - conn, dialErr := dialer.DialContext(ctx, network, net.JoinHostPort(addr.String(), port)) - if dialErr == nil { - return conn, nil - } - lastErr = dialErr - } - if lastErr == nil { - lastErr = fmt.Errorf("prompt guard no allowed resolved address") - } - return nil, lastErr - } -} - -func isExplicitPrivateHost(host string) bool { - // Only the localhost name family is trusted for private/loopback dials. - // A bare "*.local" suffix is too broad (mDNS/intranet names) and would - // re-open RFC1918 SSRF after literal private IPs were rejected. - return host == "localhost" || strings.HasSuffix(host, ".localhost") -} - -func isBlockedAddress(addr netip.Addr) bool { - if !addr.IsValid() || addr.IsUnspecified() || addr.IsMulticast() || addr.IsLinkLocalUnicast() || addr.IsLinkLocalMulticast() { - return true - } - for _, prefix := range blockedPrefixes { - if prefix.Contains(addr) { - return true - } - } - return false -} diff --git a/backend/internal/securityaudit/prompt_outbound_security_test.go b/backend/internal/securityaudit/prompt_outbound_security_test.go index 76d7f9fdb..327504d1b 100644 --- a/backend/internal/securityaudit/prompt_outbound_security_test.go +++ b/backend/internal/securityaudit/prompt_outbound_security_test.go @@ -5,7 +5,6 @@ import ( "encoding/json" "net/http" "net/http/httptest" - "net/netip" "strings" "sync/atomic" "testing" @@ -14,25 +13,20 @@ import ( "github.com/stretchr/testify/require" ) -type staticResolver struct{ addresses []netip.Addr } - -func (r staticResolver) LookupNetIP(context.Context, string, string) ([]netip.Addr, error) { - return r.addresses, nil -} - -func TestNormalizeBaseURLSecurity(t *testing.T) { - allowed := []string{"https://guard.example.com", "https://guard.example.com/v1", "http://127.0.0.1:8080", "http://localhost:8080"} +func TestNormalizeBaseURLAllowsAdministratorConfiguredDestinations(t *testing.T) { + allowed := []string{ + "https://guard.example.com", "https://guard.example.com/v1", "http://guard.example.com", + "http://127.0.0.1:8080", "http://10.0.0.8:8080", "https://172.16.0.5", + "http://169.254.169.254", "https://metadata.google.internal", "https://192.0.2.1", + "http://internal-admin.local", "http://guard.local:8080", + } for _, raw := range allowed { _, err := NormalizeBaseURL(raw) require.NoError(t, err, raw) } blocked := []string{ - "ftp://guard.example.com", "http://guard.example.com", "https://user:pass@guard.example.com", - "https://guard.example.com?q=secret", "https://guard.example.com/#fragment", "http://169.254.169.254", - "https://metadata.google.internal", "https://0.0.0.0", "https://224.0.0.1", "https://192.0.2.1", - "https://[::]", "https://[fe80::1]", "https://[ff02::1]", "https://[2001:db8::1]", - "http://10.0.0.8:8080", "http://192.168.1.10:8080", "https://172.16.0.5", - "http://internal-admin.local", "http://guard.local:8080", + "ftp://guard.example.com", "https://user:pass@guard.example.com", + "https://guard.example.com?q=secret", "https://guard.example.com/#fragment", } for _, raw := range blocked { _, err := NormalizeBaseURL(raw) @@ -43,24 +37,13 @@ func TestNormalizeBaseURLSecurity(t *testing.T) { require.Equal(t, "https://guard.example.com/v1/chat/completions", url) } -func TestSecureDialRejectsDNSRebindingToPrivateAddress(t *testing.T) { - dial := secureDialContext(nil, staticResolver{addresses: []netip.Addr{netip.MustParseAddr("127.0.0.1")}}, false) - _, err := dial(context.Background(), "tcp", "guard.example.com:443") - require.Error(t, err) -} - -func TestSecureDialLocalhostAllowlistRejectsRFC1918Resolution(t *testing.T) { - dial := secureDialContext(nil, staticResolver{addresses: []netip.Addr{netip.MustParseAddr("10.0.0.8")}}, true) - _, err := dial(context.Background(), "tcp", "localhost:8080") - require.Error(t, err) -} - -func TestSecureHTTPClientDoesNotBypassDestinationValidationThroughEnvironmentProxy(t *testing.T) { +func TestHTTPClientUsesDirectStandardDialer(t *testing.T) { client, err := NewSecureHTTPClient(ActiveEndpoint{BaseURL: "https://guard.example.com", TimeoutMS: 1000}) require.NoError(t, err) transport, ok := client.Transport.(*http.Transport) require.True(t, ok) require.Nil(t, transport.Proxy) + require.NotNil(t, transport.DialContext) } func TestOpenAICompatibleScannerRequestContract(t *testing.T) { @@ -83,13 +66,16 @@ func TestOpenAICompatibleScannerRequestContract(t *testing.T) { require.Equal(t, EventPass, result.Decision) } -func TestOpenAICompatibleScannerRejectsRedirectAndOversize(t *testing.T) { - redirect := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - http.Redirect(w, r, "http://127.0.0.1/other", http.StatusFound) +func TestOpenAICompatibleScannerFollowsRedirectAndRejectsOversize(t *testing.T) { + target := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { + _, _ = w.Write([]byte(`{"choices":[{"message":{"content":"Safety: Safe\nCategories: None"}}]}`)) })) + defer target.Close() + redirect := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { http.Redirect(w, r, target.URL, http.StatusFound) })) defer redirect.Close() - _, err := NewOpenAICompatibleScanner().Scan(context.Background(), ActiveEndpoint{ID: "redirect", BaseURL: redirect.URL, Model: DefaultGuardModel, TimeoutMS: 1000}, "hello", AllScannerIDs) - require.Error(t, err) + result, err := NewOpenAICompatibleScanner().Scan(context.Background(), ActiveEndpoint{ID: "redirect", BaseURL: redirect.URL, Model: DefaultGuardModel, TimeoutMS: 1000}, "hello", AllScannerIDs) + require.NoError(t, err) + require.Equal(t, EventPass, result.Decision) oversize := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { _, _ = w.Write([]byte(strings.Repeat("x", int(maxGuardResponseBytes)+1))) })) diff --git a/deploy/build_image.sh b/deploy/build_image.sh old mode 100644 new mode 100755 diff --git a/deploy/docker-compose.yml b/deploy/docker-compose.yml index 6aecdcfa5..637c2762f 100644 --- a/deploy/docker-compose.yml +++ b/deploy/docker-compose.yml @@ -16,7 +16,7 @@ services: # Sub2API Application # =========================================================================== sub2api: - image: weishaw/sub2api:latest + image: sub2api:latest container_name: sub2api restart: unless-stopped ulimits: diff --git a/frontend/src/features/prompt-audit/PromptAuditView.vue b/frontend/src/features/prompt-audit/PromptAuditView.vue index e0491bbdc..939fc52ac 100644 --- a/frontend/src/features/prompt-audit/PromptAuditView.vue +++ b/frontend/src/features/prompt-audit/PromptAuditView.vue @@ -1,6 +1,6 @@