Skip to content

Commit 4b5f1ea

Browse files
committed
feat(plugin): support observing upstream websocket response events
- Introduce `WebSocketResponseObserver` capability and bump plugin ABI schema version to 4. - Forward upstream WebSocket response frames from Codex and xAI executors to configured observers. - Wire `WebSocketResponseObserver` across API handlers and plugin host dispatchers. Closes: router-for-me#5248
1 parent b7f6c15 commit 4b5f1ea

20 files changed

Lines changed: 810 additions & 16 deletions

internal/pluginhost/adapters_interceptors.go

Lines changed: 55 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -72,6 +72,20 @@ func (h *Host) callStreamChunkInterceptor(ctx context.Context, record capability
7272
return resp, true
7373
}
7474

75+
func (h *Host) callWebSocketResponseObserver(ctx context.Context, record capabilityRecord, observer pluginapi.WebSocketResponseObserver, event pluginapi.WebSocketResponseEvent) {
76+
if h == nil || observer == nil || h.isPluginFused(record.id) || !h.recordCurrent(record) {
77+
return
78+
}
79+
defer func() {
80+
if recovered := recover(); recovered != nil {
81+
h.fusePlugin(record.id, "WebSocketResponseObserver.ObserveWebSocketResponseEvent", recovered)
82+
}
83+
}()
84+
if errObserve := observer.ObserveWebSocketResponseEvent(ctx, event); errObserve != nil {
85+
log.Warnf("pluginhost: websocket response observer %s failed: %v", record.id, errObserve)
86+
}
87+
}
88+
7589
func (h *Host) InterceptRequestBeforeAuth(ctx context.Context, req pluginapi.RequestInterceptRequest) pluginapi.RequestInterceptResponse {
7690
return h.InterceptRequestBeforeAuthExcept(ctx, req, "")
7791
}
@@ -131,6 +145,47 @@ func (h *Host) CompleteRequest(ctx context.Context, completion pluginapi.Request
131145
h.CompleteRequestExcept(ctx, completion, "")
132146
}
133147

148+
func (h *Host) HasWebSocketResponseObservers() bool {
149+
if h == nil {
150+
return false
151+
}
152+
for _, record := range h.activeRecords() {
153+
if h.isPluginFused(record.id) {
154+
continue
155+
}
156+
if record.plugin.Capabilities.WebSocketResponseObserver != nil {
157+
return true
158+
}
159+
}
160+
return false
161+
}
162+
163+
func (h *Host) ObserveWebSocketResponseEvent(ctx context.Context, event pluginapi.WebSocketResponseEvent) {
164+
h.ObserveWebSocketResponseEventExcept(ctx, event, "")
165+
}
166+
167+
func (h *Host) ObserveWebSocketResponseEventExcept(ctx context.Context, event pluginapi.WebSocketResponseEvent, skipPluginID string) {
168+
if h == nil {
169+
return
170+
}
171+
if ctx == nil {
172+
ctx = context.Background()
173+
} else {
174+
ctx = context.WithoutCancel(ctx)
175+
}
176+
skipPluginID = strings.TrimSpace(skipPluginID)
177+
for _, record := range h.activeRecords() {
178+
observer := record.plugin.Capabilities.WebSocketResponseObserver
179+
if h.isPluginFused(record.id) || observer == nil || record.id == skipPluginID || !h.recordCurrent(record) {
180+
continue
181+
}
182+
next := event
183+
next.Payload = bytes.Clone(event.Payload)
184+
next.Metadata = cloneInterceptorMetadata(event.Metadata)
185+
h.callWebSocketResponseObserver(ctx, record, observer, next)
186+
}
187+
}
188+
134189
// CompleteRequestExcept notifies lifecycle plugins except the plugin that initiated a nested host execution.
135190
func (h *Host) CompleteRequestExcept(ctx context.Context, completion pluginapi.RequestCompletion, skipPluginID string) {
136191
if h == nil {

internal/pluginhost/host.go

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -1052,6 +1052,7 @@ func validPlugin(plugin pluginapi.Plugin) bool {
10521052
caps.ResponseAfterTranslator != nil ||
10531053
caps.ResponseInterceptor != nil ||
10541054
caps.StreamChunkInterceptor != nil ||
1055+
caps.WebSocketResponseObserver != nil ||
10551056
caps.ThinkingApplier != nil ||
10561057
caps.UsagePlugin != nil ||
10571058
caps.CommandLinePlugin != nil ||

internal/pluginhost/rpc_client.go

Lines changed: 19 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -133,6 +133,9 @@ func registerRPCPlugin(ctx context.Context, host *Host, id string, client plugin
133133
if resp.Capabilities.StreamChunkInterceptor {
134134
plugin.Capabilities.StreamChunkInterceptor = adapter
135135
}
136+
if resp.Capabilities.WebSocketResponseObserver {
137+
plugin.Capabilities.WebSocketResponseObserver = adapter
138+
}
136139
if resp.Capabilities.ThinkingApplier {
137140
plugin.Capabilities.ThinkingApplier = rpcThinkingApplier{rpcPluginAdapter: adapter}
138141
}
@@ -211,6 +214,9 @@ func sanitizePluginRequest(request any) any {
211214
case pluginapi.StreamChunkInterceptRequest:
212215
req.Metadata = sanitizePluginMetadata(req.Metadata)
213216
return req
217+
case pluginapi.WebSocketResponseEvent:
218+
req.Metadata = sanitizePluginMetadata(req.Metadata)
219+
return req
214220
case rpcRequestInterceptRequest:
215221
req.Metadata = sanitizePluginMetadata(req.Metadata)
216222
return req
@@ -226,6 +232,9 @@ func sanitizePluginRequest(request any) any {
226232
case rpcStreamChunkInterceptRequest:
227233
req.Metadata = sanitizePluginMetadata(req.Metadata)
228234
return req
235+
case rpcWebSocketResponseEvent:
236+
req.Metadata = sanitizePluginMetadata(req.Metadata)
237+
return req
229238
case pluginapi.ExecutorHTTPRequest:
230239
req.HTTPClient = nil
231240
return req
@@ -530,6 +539,16 @@ func (a *rpcPluginAdapter) InterceptStreamChunk(ctx context.Context, req plugina
530539
})
531540
}
532541

542+
func (a *rpcPluginAdapter) ObserveWebSocketResponseEvent(ctx context.Context, event pluginapi.WebSocketResponseEvent) error {
543+
callbackID, closeCallback := a.openHostCallbackContext(ctx)
544+
defer closeCallback()
545+
_, errCall := callPlugin[rpcEmptyResponse](ctx, a.client, pluginabi.MethodWebSocketResponseEvent, rpcWebSocketResponseEvent{
546+
WebSocketResponseEvent: event,
547+
HostCallbackID: callbackID,
548+
})
549+
return errCall
550+
}
551+
533552
func (a rpcThinkingApplier) ApplyThinking(ctx context.Context, req pluginapi.ThinkingApplyRequest) (pluginapi.PayloadResponse, error) {
534553
callbackID, closeCallback := a.openHostCallbackContext(ctx)
535554
defer closeCallback()

internal/pluginhost/rpc_schema.go

Lines changed: 7 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -39,6 +39,7 @@ type rpcCapabilities struct {
3939
ResponseAfterTranslator bool `json:"response_after_translator"`
4040
ResponseInterceptor bool `json:"response_interceptor"`
4141
StreamChunkInterceptor bool `json:"response_stream_interceptor"`
42+
WebSocketResponseObserver bool `json:"websocket_response_observer"`
4243
ThinkingApplier bool `json:"thinking_applier"`
4344
UsagePlugin bool `json:"usage_plugin"`
4445
CommandLinePlugin bool `json:"command_line_plugin"`
@@ -110,6 +111,11 @@ type rpcStreamChunkInterceptRequest struct {
110111
HostCallbackID string `json:"host_callback_id,omitempty"`
111112
}
112113

114+
type rpcWebSocketResponseEvent struct {
115+
pluginapi.WebSocketResponseEvent
116+
HostCallbackID string `json:"host_callback_id,omitempty"`
117+
}
118+
113119
type rpcThinkingApplyRequest struct {
114120
pluginapi.ThinkingApplyRequest
115121
HostCallbackID string `json:"host_callback_id,omitempty"`
@@ -150,6 +156,7 @@ func rpcCapabilitiesFromPlugin(plugin pluginapi.Plugin) rpcCapabilities {
150156
ResponseAfterTranslator: caps.ResponseAfterTranslator != nil,
151157
ResponseInterceptor: caps.ResponseInterceptor != nil,
152158
StreamChunkInterceptor: caps.StreamChunkInterceptor != nil,
159+
WebSocketResponseObserver: caps.WebSocketResponseObserver != nil,
153160
ThinkingApplier: caps.ThinkingApplier != nil,
154161
UsagePlugin: caps.UsagePlugin != nil,
155162
CommandLinePlugin: caps.CommandLinePlugin != nil,
Lines changed: 247 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,247 @@
1+
package pluginhost
2+
3+
import (
4+
"bytes"
5+
"context"
6+
"encoding/json"
7+
"testing"
8+
9+
"github.com/router-for-me/CLIProxyAPI/v7/sdk/pluginabi"
10+
"github.com/router-for-me/CLIProxyAPI/v7/sdk/pluginapi"
11+
)
12+
13+
type testWebSocketObserverFunc func(context.Context, pluginapi.WebSocketResponseEvent) error
14+
15+
func (f testWebSocketObserverFunc) ObserveWebSocketResponseEvent(ctx context.Context, event pluginapi.WebSocketResponseEvent) error {
16+
return f(ctx, event)
17+
}
18+
19+
func TestObserveWebSocketResponseEventInvokesPlugin(t *testing.T) {
20+
var got pluginapi.WebSocketResponseEvent
21+
called := false
22+
host := newHostWithRecords(capabilityRecord{
23+
id: "quota-tracker",
24+
plugin: pluginapi.Plugin{
25+
Capabilities: pluginapi.Capabilities{
26+
WebSocketResponseObserver: testWebSocketObserverFunc(func(_ context.Context, event pluginapi.WebSocketResponseEvent) error {
27+
got = event
28+
called = true
29+
return nil
30+
}),
31+
},
32+
},
33+
})
34+
35+
rawPayload := []byte(`{"type":"codex.rate_limits","rate_limits":{"primary":{"used_percent":42}}}`)
36+
host.ObserveWebSocketResponseEvent(context.Background(), pluginapi.WebSocketResponseEvent{
37+
RequestID: "req-123",
38+
SourceFormat: "openai",
39+
Model: "gpt-5.3-codex",
40+
RequestedModel: "gpt-5.3-codex",
41+
Provider: "codex",
42+
AuthID: "auth-abc",
43+
AuthLabel: "test-auth",
44+
AuthType: "oauth",
45+
EventType: "codex.rate_limits",
46+
Payload: rawPayload,
47+
})
48+
49+
if !called {
50+
t.Fatal("observer callback was not invoked")
51+
}
52+
if got.RequestID != "req-123" {
53+
t.Fatalf("RequestID = %q, want req-123", got.RequestID)
54+
}
55+
if got.Provider != "codex" {
56+
t.Fatalf("Provider = %q, want codex", got.Provider)
57+
}
58+
if got.AuthID != "auth-abc" || got.AuthLabel != "test-auth" {
59+
t.Fatalf("Auth = (%q, %q), want (auth-abc, test-auth)", got.AuthID, got.AuthLabel)
60+
}
61+
if got.EventType != "codex.rate_limits" {
62+
t.Fatalf("EventType = %q, want codex.rate_limits", got.EventType)
63+
}
64+
if !bytes.Equal(got.Payload, rawPayload) {
65+
t.Fatalf("Payload = %s, want %s", got.Payload, rawPayload)
66+
}
67+
}
68+
69+
func TestObserveWebSocketResponseEventClonesPayloadAndMetadata(t *testing.T) {
70+
ctx, cancel := context.WithCancel(context.Background())
71+
cancel()
72+
originalPayload := []byte(`{"type":"codex.rate_limits","rate_limits":{"primary":{"used_percent":1}}}`)
73+
originalMetadata := map[string]any{"key": "value"}
74+
called := false
75+
76+
host := newHostWithRecords(capabilityRecord{
77+
id: "observer-clone",
78+
plugin: pluginapi.Plugin{
79+
Capabilities: pluginapi.Capabilities{
80+
WebSocketResponseObserver: testWebSocketObserverFunc(func(_ context.Context, event pluginapi.WebSocketResponseEvent) error {
81+
event.Payload[0] = 'X'
82+
event.Metadata["key"] = "mutated"
83+
called = true
84+
return nil
85+
}),
86+
},
87+
},
88+
})
89+
90+
host.ObserveWebSocketResponseEvent(ctx, pluginapi.WebSocketResponseEvent{
91+
RequestID: "req-clone",
92+
Payload: originalPayload,
93+
Metadata: originalMetadata,
94+
})
95+
96+
if !called {
97+
t.Fatal("observer callback was not invoked")
98+
}
99+
if originalPayload[0] != '{' {
100+
t.Fatalf("original payload was mutated: %s", originalPayload)
101+
}
102+
if originalMetadata["key"] != "value" {
103+
t.Fatalf("original metadata was mutated: %#v", originalMetadata)
104+
}
105+
}
106+
107+
func TestObserveWebSocketResponseEventFusesOnPanic(t *testing.T) {
108+
host := newHostWithRecords(capabilityRecord{
109+
id: "panicking-observer",
110+
plugin: pluginapi.Plugin{
111+
Capabilities: pluginapi.Capabilities{
112+
WebSocketResponseObserver: testWebSocketObserverFunc(func(context.Context, pluginapi.WebSocketResponseEvent) error {
113+
panic("observer panic")
114+
}),
115+
},
116+
},
117+
})
118+
119+
if !host.HasWebSocketResponseObservers() {
120+
t.Fatal("HasWebSocketResponseObservers() = false, want true")
121+
}
122+
123+
host.ObserveWebSocketResponseEvent(context.Background(), pluginapi.WebSocketResponseEvent{
124+
RequestID: "req-panic",
125+
Payload: []byte(`{"type":"test"}`),
126+
})
127+
128+
if !host.isPluginFused("panicking-observer") {
129+
t.Fatal("isPluginFused(panicking-observer) = false, want true")
130+
}
131+
if host.HasWebSocketResponseObservers() {
132+
t.Fatal("HasWebSocketResponseObservers() = true after fusing, want false")
133+
}
134+
}
135+
136+
func TestObserveWebSocketResponseEventSkipPlugin(t *testing.T) {
137+
calls := 0
138+
host := newHostWithRecords(capabilityRecord{
139+
id: "skipped-plugin",
140+
plugin: pluginapi.Plugin{
141+
Capabilities: pluginapi.Capabilities{
142+
WebSocketResponseObserver: testWebSocketObserverFunc(func(context.Context, pluginapi.WebSocketResponseEvent) error {
143+
calls++
144+
return nil
145+
}),
146+
},
147+
},
148+
})
149+
150+
host.ObserveWebSocketResponseEventExcept(context.Background(), pluginapi.WebSocketResponseEvent{
151+
RequestID: "req-skip",
152+
}, "skipped-plugin")
153+
154+
if calls != 0 {
155+
t.Fatalf("observer calls = %d, want 0 when skipped", calls)
156+
}
157+
}
158+
159+
func TestRegisterRPCPluginRegistersWebSocketResponseObserver(t *testing.T) {
160+
lookup := newTestSymbolLookup(&testPlugin{
161+
registerResult: pluginapi.Plugin{
162+
Capabilities: pluginapi.Capabilities{
163+
WebSocketResponseObserver: testWebSocketObserverFunc(func(context.Context, pluginapi.WebSocketResponseEvent) error {
164+
return nil
165+
}),
166+
},
167+
},
168+
})
169+
170+
registered, errRegister := registerRPCPlugin(context.Background(), nil, "rpc-observer", lookup, pluginabi.MethodPluginRegister, nil)
171+
if errRegister != nil {
172+
t.Fatalf("registerRPCPlugin() error = %v", errRegister)
173+
}
174+
if registered.Capabilities.WebSocketResponseObserver == nil {
175+
t.Fatal("WebSocketResponseObserver = nil, want RPC adapter")
176+
}
177+
}
178+
179+
type rpcObserverRecordingClient struct {
180+
lastMethod string
181+
lastRequest []byte
182+
}
183+
184+
func (c *rpcObserverRecordingClient) Call(_ context.Context, method string, request []byte) ([]byte, error) {
185+
c.lastMethod = method
186+
c.lastRequest = bytes.Clone(request)
187+
return json.Marshal(pluginabi.Envelope{OK: true, Result: json.RawMessage(`{}`)})
188+
}
189+
190+
func (c *rpcObserverRecordingClient) Shutdown() {}
191+
192+
func TestObserveWebSocketResponseEventRPCSanitizesMetadata(t *testing.T) {
193+
client := &rpcObserverRecordingClient{}
194+
adapter := &rpcPluginAdapter{
195+
id: "rpc-observer-sanitize",
196+
client: client,
197+
}
198+
199+
rawPayload := []byte(`{"type":"codex.rate_limits","rate_limits":{"primary":{"used_percent":99}}}`)
200+
unserializableMetadata := map[string]any{
201+
"safe_key": "safe_value",
202+
"func_field": func() {},
203+
"chan_field": make(chan int),
204+
}
205+
206+
err := adapter.ObserveWebSocketResponseEvent(context.Background(), pluginapi.WebSocketResponseEvent{
207+
RequestID: "req-rpc-sanitize",
208+
SourceFormat: "openai",
209+
Model: "gpt-5.3-codex",
210+
RequestedModel: "gpt-5.3-codex",
211+
Provider: "codex",
212+
AuthID: "auth-xyz",
213+
AuthLabel: "test-auth-rpc",
214+
AuthType: "oauth",
215+
EventType: "codex.rate_limits",
216+
Payload: rawPayload,
217+
Metadata: unserializableMetadata,
218+
})
219+
if err != nil {
220+
t.Fatalf("ObserveWebSocketResponseEvent() error = %v", err)
221+
}
222+
223+
if client.lastMethod != pluginabi.MethodWebSocketResponseEvent {
224+
t.Fatalf("lastMethod = %q, want %q", client.lastMethod, pluginabi.MethodWebSocketResponseEvent)
225+
}
226+
227+
var decoded rpcWebSocketResponseEvent
228+
if errUnmarshal := json.Unmarshal(client.lastRequest, &decoded); errUnmarshal != nil {
229+
t.Fatalf("unmarshal rpc request: %v", errUnmarshal)
230+
}
231+
232+
if decoded.RequestID != "req-rpc-sanitize" {
233+
t.Fatalf("decoded RequestID = %q, want req-rpc-sanitize", decoded.RequestID)
234+
}
235+
if decoded.EventType != "codex.rate_limits" {
236+
t.Fatalf("decoded EventType = %q, want codex.rate_limits", decoded.EventType)
237+
}
238+
if decoded.Metadata["safe_key"] != "safe_value" {
239+
t.Fatalf("decoded safe_key = %v, want safe_value", decoded.Metadata["safe_key"])
240+
}
241+
if _, exists := decoded.Metadata["func_field"]; exists {
242+
t.Fatalf("func_field was not sanitized: %#v", decoded.Metadata)
243+
}
244+
if _, exists := decoded.Metadata["chan_field"]; exists {
245+
t.Fatalf("chan_field was not sanitized: %#v", decoded.Metadata)
246+
}
247+
}

0 commit comments

Comments
 (0)