Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
3 changes: 3 additions & 0 deletions .gitignore
Original file line number Diff line number Diff line change
Expand Up @@ -49,3 +49,6 @@ workspace.zip
/.trellis
/AGENTS.md
/CLAUDE.md

# 本地工具链(protoc/LLVM),不提交
/.tools/
7 changes: 7 additions & 0 deletions crates/agent-gateway/internal/handler/image_proxy.go
Original file line number Diff line number Diff line change
Expand Up @@ -79,6 +79,13 @@ func imageProxyWithClient(client outboundHTTPClient) http.HandlerFunc {
w.Header().Set("Cache-Control", "private, max-age=300")
w.Header().Set("X-Content-Type-Options", "nosniff")
w.Header().Set("Referrer-Policy", "no-referrer")
if mimeType == "image/svg+xml" {
// SVG 可携带脚本;顶层导航到 SVG 会在网关 origin 执行脚本(WebUI
// 的管理 token 在 localStorage)。sandbox 禁脚本 + 附件式内联处置
// 把该端点降级为纯图片源。
w.Header().Set("Content-Security-Policy", "sandbox; default-src 'none'")
w.Header().Set("Content-Disposition", `inline; filename="image.svg"`)
}
_, _ = w.Write(body)
}
}
Expand Down
56 changes: 56 additions & 0 deletions crates/agent-gateway/internal/handler/image_proxy_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -46,6 +46,62 @@ func TestImageProxyServesSupportedImage(t *testing.T) {
}
}

func TestImageProxyServesSVGWithSandboxCSP(t *testing.T) {
// 顶层导航到 SVG 会在网关 origin 执行脚本(WebUI 管理 token 在 localStorage);
// SVG 响应必须携带 sandbox CSP 与附件式处置,降级为纯图片源。
client := outboundHTTPClientFunc(func(r *http.Request) (*http.Response, error) {
body := []byte(`<svg xmlns="http://www.w3.org/2000/svg"><script>alert(1)</script></svg>`)
return &http.Response{
StatusCode: http.StatusOK,
Header: http.Header{"Content-Type": []string{"text/plain"}},
Body: io.NopCloser(strings.NewReader(string(body))),
ContentLength: int64(len(body)),
Request: r,
}, nil
})

req := httptest.NewRequest(http.MethodGet, "/image-proxy?url=https://evil.example/x.svg", nil)
rec := httptest.NewRecorder()
imageProxyWithClient(client)(rec, req)

if rec.Code != http.StatusOK {
t.Fatalf("expected status %d, got %d body=%q", http.StatusOK, rec.Code, rec.Body.String())
}
if got := rec.Header().Get("Content-Type"); got != "image/svg+xml" {
t.Fatalf("content-type = %q, want image/svg+xml", got)
}
if got := rec.Header().Get("Content-Security-Policy"); got != "sandbox; default-src 'none'" {
t.Fatalf("svg CSP = %q, want sandbox CSP", got)
}
if got := rec.Header().Get("Content-Disposition"); !strings.Contains(got, "inline") {
t.Fatalf("svg content-disposition = %q, want inline attachment-style disposition", got)
}
}

func TestImageProxyNonSVGHasNoSandboxCSP(t *testing.T) {
client := outboundHTTPClientFunc(func(r *http.Request) (*http.Response, error) {
body := []byte("\x89PNG\r\n\x1a\nliveagent-test")
return &http.Response{
StatusCode: http.StatusOK,
Header: http.Header{"Content-Type": []string{"image/png"}},
Body: io.NopCloser(strings.NewReader(string(body))),
ContentLength: int64(len(body)),
Request: r,
}, nil
})

req := httptest.NewRequest(http.MethodGet, "/image-proxy?url=https://images.example/photo.png", nil)
rec := httptest.NewRecorder()
imageProxyWithClient(client)(rec, req)

if rec.Code != http.StatusOK {
t.Fatalf("expected status %d, got %d body=%q", http.StatusOK, rec.Code, rec.Body.String())
}
if got := rec.Header().Get("Content-Security-Policy"); got != "" {
t.Fatalf("non-svg CSP = %q, want empty", got)
}
}

func TestImageProxyRefererUsesTargetOrigin(t *testing.T) {
targetURL, err := url.Parse("https://example.com:8443/path/photo.png?size=large")
if err != nil {
Expand Down
43 changes: 27 additions & 16 deletions crates/agent-gateway/internal/protocol/pbws/browser_conn.go
Original file line number Diff line number Diff line change
Expand Up @@ -123,18 +123,16 @@ func (c *browserConn) serve() {
defer observability.Usage.V2BrowserConnectionsActive.Add(-1)

for {
frame, ok := c.readFrame()
frame, denied, ok := c.readFrame()
if !ok {
return
}
c.core.TouchInboundActivity()

// 入站限速:超限丢帧回错,连续违规判定失控客户端、关闭连接。
if allowed, exceeded := c.rateLimiter.Allow(); !allowed {
if exceeded {
return
}
_ = c.sendLocalError(frame.GetRequestId(), "too many requests")
// 入站限速已在 readFrame 内、反序列化前计数(文本帧同罪);被丢弃的帧
// 无 request 上下文,回无关联 id 的本地错误,连续违规由 readFrame 判死。
if denied {
_ = c.sendLocalError("", "too many requests")
continue
}

Expand Down Expand Up @@ -168,29 +166,42 @@ func (c *browserConn) serve() {
}

// readFrame 读取并解码一帧;解码失败说明帧流已破坏,直接关闭连接。
func (c *browserConn) readFrame() (*gatewayv2.WebClientFrame, bool) {
// 限速在每次 ReadMessage 成功后、proto 反序列化前立即计数:文本帧与大型
// 二进制帧同样消耗令牌——旧实现把 Allow 放在反序列化之后且文本帧直接
// continue,洪泛文本帧可零成本绕过限速,大帧的解析 CPU 也无法被约束。
// denied=true 表示该帧在解析前被限速丢弃、连接仍存活;ok=false 表示连接
// 应关闭(读错误 / 解码失败 / 连续违规超阈值判定失控客户端)。
func (c *browserConn) readFrame() (frame *gatewayv2.WebClientFrame, denied bool, ok bool) {
for {
messageType, data, err := c.conn.ReadMessage()
if err != nil {
return nil, false
return nil, false, false
}
if allowed, exceeded := c.rateLimiter.Allow(); !allowed {
if exceeded {
return nil, false, false
}
return nil, true, true
}
if messageType != websocket.BinaryMessage {
// v2 链路上文本帧无意义;容忍并忽略(仍计入存活)。
// v2 链路上文本帧无意义;计入限速(上方)后忽略(仍计入存活)。
c.core.TouchInboundActivity()
continue
}
var frame gatewayv2.WebClientFrame
if err := proto.Unmarshal(data, &frame); err != nil {
return nil, false
var decoded gatewayv2.WebClientFrame
if err := proto.Unmarshal(data, &decoded); err != nil {
return nil, false, false
}
return &frame, true
return &decoded, false, true
}
}

// handshake 处理首帧 hello;失败时写出失败应答并关闭。
func (c *browserConn) handshake() bool {
frame, ok := c.readFrame()
if !ok {
// 首帧即被限速(前置洪泛耗尽突发额度)说明对端行为异常,直接关闭;
// 正常客户端的 hello 是第一帧,不会命中限速。
frame, denied, ok := c.readFrame()
if !ok || denied {
return false
}
hello := frame.GetHello()
Expand Down
128 changes: 128 additions & 0 deletions crates/agent-gateway/internal/protocol/pbws/browser_conn_test.go
Original file line number Diff line number Diff line change
@@ -0,0 +1,128 @@
package pbws

// 主链路读路径与终端链路同形态的限速绑定回归:文本帧计入限速、限速先于
// 反序列化(旧实现 Allow 在 Unmarshal 之后且文本帧直接 continue,可被
// 零成本文本帧洪泛绕过,大帧的解析 CPU 也无法被约束)。

import (
"testing"
"time"

"github.com/gorilla/websocket"
"google.golang.org/protobuf/proto"

gatewayv2 "github.com/liveagent/agent-gateway/internal/proto/v2"
"github.com/liveagent/agent-gateway/internal/transport/wscore"
)

// newReadFrameTestConn 构造仅够驱动 readFrame 的轻量 browserConn:
// 读路径只触达 conn / rateLimiter / core.TouchInboundActivity。
func newReadFrameTestConn(conn *websocket.Conn, limiter *wscore.InboundRateLimiter) *browserConn {
return &browserConn{
conn: conn,
core: wscore.NewConn(conn, wscore.Config{}),
rateLimiter: limiter,
}
}

func mustMarshalWebClientFrame(t *testing.T) []byte {
t.Helper()
data, err := proto.Marshal(&gatewayv2.WebClientFrame{})
if err != nil {
t.Fatalf("marshal web client frame: %v", err)
}
return data
}

// 文本帧必须与二进制帧同罪计入限速:两帧文本耗尽突发额度后,合法二进制帧
// 必须被限速丢弃;旧实现中文本帧不计数,该帧会被放行。
func TestBrowserReadFrameCountsTextFramesAgainstRateLimit(t *testing.T) {
results := make(chan frameReadResult, 3)
limiter := wscore.NewInboundRateLimiter(0.0001, 2, 3)
client := wsTestPair(t, func(conn *websocket.Conn) {
defer func() { _ = conn.Close() }()
c := newReadFrameTestConn(conn, limiter)
for i := 0; i < 3; i++ {
frame, denied, ok := c.readFrame()
results <- frameReadResult{gotFrame: frame != nil, denied: denied, ok: ok}
if !ok {
return
}
}
})

valid := mustMarshalWebClientFrame(t)
writes := []struct {
messageType int
data []byte
}{
{websocket.TextMessage, []byte("junk-one")},
{websocket.TextMessage, []byte("junk-two")},
{websocket.BinaryMessage, valid},
{websocket.BinaryMessage, valid},
{websocket.BinaryMessage, valid},
}
for _, write := range writes {
if err := client.WriteMessage(write.messageType, write.data); err != nil {
t.Fatalf("write frame: %v", err)
}
}

first := <-results
if !first.denied || !first.ok || first.gotFrame {
t.Fatalf("first read = %+v, want denied-but-alive (text frames must consume budget)", first)
}
second := <-results
if !second.denied || !second.ok || second.gotFrame {
t.Fatalf("second read = %+v, want denied-but-alive", second)
}
third := <-results
if third.ok {
t.Fatalf("third read = %+v, want connection closed after 3 consecutive violations", third)
}
}

// 限速必须发生在 proto 反序列化之前:超限的非法帧在解析前即被丢弃、连接
// 存活;旧顺序(先解析后限速)下非法帧直接破坏帧流、关闭连接,第三帧
// 永远读不到。
func TestBrowserReadFrameLimitsBeforeUnmarshal(t *testing.T) {
results := make(chan frameReadResult, 3)
limiter := wscore.NewInboundRateLimiter(100, 1, 10)
client := wsTestPair(t, func(conn *websocket.Conn) {
defer func() { _ = conn.Close() }()
c := newReadFrameTestConn(conn, limiter)
for i := 0; i < 3; i++ {
frame, denied, ok := c.readFrame()
results <- frameReadResult{gotFrame: frame != nil, denied: denied, ok: ok}
if !ok {
return
}
}
})

valid := mustMarshalWebClientFrame(t)
if err := client.WriteMessage(websocket.BinaryMessage, valid); err != nil {
t.Fatalf("write valid frame: %v", err)
}
if err := client.WriteMessage(websocket.BinaryMessage, []byte{0xff, 0xff, 0xff}); err != nil {
t.Fatalf("write invalid frame: %v", err)
}
// 等限速器回填后第三帧必须仍然可读、可解析。
time.Sleep(30 * time.Millisecond)
if err := client.WriteMessage(websocket.BinaryMessage, valid); err != nil {
t.Fatalf("write third frame: %v", err)
}

first := <-results
if !first.gotFrame || !first.ok || first.denied {
t.Fatalf("first read = %+v, want parsed frame", first)
}
second := <-results
if !second.denied || !second.ok || second.gotFrame {
t.Fatalf("second read = %+v, want denied-but-alive (rate limit must precede unmarshal)", second)
}
third := <-results
if !third.gotFrame || !third.ok || third.denied {
t.Fatalf("third read = %+v, want parsed frame after refill", third)
}
}
4 changes: 3 additions & 1 deletion crates/agent-gateway/internal/protocol/pbws/guard.go
Original file line number Diff line number Diff line change
Expand Up @@ -73,7 +73,9 @@ func vetAgentRequest(sm session.AgentView, env *gatewayv2.GatewayEnvelope) error

// ---- 带功能门控 / 限额的直通臂 ----
case *gatewayv2.GatewayEnvelope_GitRequest:
action := strings.TrimSpace(payload.GitRequest.GetAction())
// 大小写归一:桌面侧对 action 做 to_ascii_lowercase 后执行,若此处不归一,
// CLONE_START/Push 等变体可绕过写操作门控(判为读操作)却在桌面端执行真实写操作。
action := strings.ToLower(strings.TrimSpace(payload.GitRequest.GetAction()))
if gitActionIsWrite(action) && !sm.WebGitEnabled() {
return errors.New("web git is disabled in desktop Remote settings")
}
Expand Down
Loading
Loading