Skip to content
Merged
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
6 changes: 6 additions & 0 deletions README.md
Original file line number Diff line number Diff line change
Expand Up @@ -346,6 +346,12 @@ height; request payloads and upstream error messages are omitted.

### Subscription resume

Send `eth_subscribe` and `eth_unsubscribe` as individual WebSocket requests.
Stitch rejects any batch containing either method before forwarding any member,
including ordinary calls in a mixed batch. Valid requests receive a `-32000`
error with their original IDs in a batch response; notifications receive no
reply. Ordinary RPC batches remain supported.

When a client opens an `eth_subscribe newHeads`, stitch:

1. Mints a synthetic subscription ID and returns it to the client.
Expand Down
72 changes: 72 additions & 0 deletions internal/subscription/eth_batch.go
Original file line number Diff line number Diff line change
@@ -0,0 +1,72 @@
package subscription

import (
"bytes"
"encoding/json"

"github.com/InjectiveLabs/stitch/internal/cache"
)

// rejectEthSubscriptionBatch preflights the whole batch before any ID allocation
// or upstream write. Batch subscription acknowledgments cannot use the ordinary
// RPC correlation path: their IDs must belong to a tracked, resumable subscription.
// Reject mixed batches atomically so ordinary members cannot execute silently.
func rejectEthSubscriptionBatch(io sessionIO, msg []byte) (bool, error) {
trimmed := bytes.TrimSpace(msg)
if len(trimmed) == 0 || trimmed[0] != '[' {
return false, nil
}
var batch []json.RawMessage
if err := json.Unmarshal(trimmed, &batch); err != nil {
return false, nil // Preserve upstream handling of malformed JSON.
}
members := make([]map[string]json.RawMessage, len(batch))
containsSubscription := false
for i, raw := range batch {
_ = json.Unmarshal(raw, &members[i])
var method string
_ = json.Unmarshal(members[i]["method"], &method)
if method == "eth_subscribe" || method == "eth_unsubscribe" {
containsSubscription = true
}
}
if !containsSubscription {
return false, nil
}

responses := make([]map[string]any, 0, len(batch))
for _, member := range members {
var method, version string
methodJSON := bytes.TrimSpace(member["method"])
valid := len(methodJSON) > 0 && methodJSON[0] == '"' &&
json.Unmarshal(methodJSON, &method) == nil &&
json.Unmarshal(member["jsonrpc"], &version) == nil && version == "2.0"
if params, exists := member["params"]; exists {
params = bytes.TrimSpace(params)
valid = valid && len(params) > 0 && (params[0] == '[' || params[0] == '{')
}
id, hasID := member["id"]
valid = valid && (!hasID || cache.IsJSONRPCID(id))
if valid && !hasID {
continue // Notifications never receive a response, even on rejection.
}
code := -32000
message := "batch contains subscription methods; send requests individually"
if !valid {
id, code, message = nil, -32600, "invalid request"
}
responses = append(responses, map[string]any{
"jsonrpc": "2.0",
"id": id,
"error": map[string]any{"code": code, "message": message},
})
}
if len(responses) == 0 {
return true, nil
}
out, err := json.Marshal(responses)
if err != nil {
return true, err
}
return true, io.clientWrite(out)
}
95 changes: 95 additions & 0 deletions internal/subscription/eth_batch_test.go
Original file line number Diff line number Diff line change
@@ -0,0 +1,95 @@
package subscription

import (
"encoding/json"
"testing"
)

func TestEthSubscriptionBatchesAreRejectedBeforeForwarding(t *testing.T) {
for _, tc := range []struct {
name string
body string
ids []string
codes []int
}{
{"subscribe", `[{"jsonrpc":"2.0","id":1,"method":"eth_subscribe","params":["newHeads"]}]`, []string{"1"}, []int{-32000}},
{"mixed IDs", `[{"jsonrpc":"2.0","id":9007199254740993,"method":"eth_chainId"},{"jsonrpc":"2.0","id":"1","method":"eth_subscribe","params":["newHeads"]},{"jsonrpc":"2.0","id":null,"method":"eth_chainId"},{"jsonrpc":"2.0","id":"1","method":"eth_chainId"}]`, []string{"9007199254740993", `"1"`, "null", `"1"`}, []int{-32000, -32000, -32000, -32000}},
{"unsubscribe", `[{"jsonrpc":"2.0","id":"cancel","method":"eth_unsubscribe","params":["0x0000000000000001"]},{"jsonrpc":"2.0","id":2,"method":"eth_chainId"}]`, []string{`"cancel"`, "2"}, []int{-32000, -32000}},
{"notifications only", `[{"jsonrpc":"2.0","method":"eth_subscribe","params":["newHeads"]},{"jsonrpc":"2.0","method":"eth_chainId"}]`, nil, nil},
{"mixed notification", `[{"jsonrpc":"2.0","id":2,"method":"eth_chainId"},{"jsonrpc":"2.0","method":"eth_unsubscribe","params":["0x0000000000000001"]}]`, []string{"2"}, []int{-32000}},
{"invalid members", `[1,null,{"foo":"bar"},{"jsonrpc":"2.0","id":{},"method":"eth_chainId"},{"jsonrpc":"2.0","id":3,"method":"eth_subscribe","params":["newHeads"]}]`, []string{"null", "null", "null", "null", "3"}, []int{-32600, -32600, -32600, -32600, -32000}},
{"invalid notification and subscription", `[{"jsonrpc":"2.0","method":null},{"jsonrpc":"2.0","id":4,"method":null},{"jsonrpc":1,"id":9,"method":"eth_subscribe"},{"jsonrpc":"2.0","method":"eth_chainId","params":true}]`, []string{"null", "null", "null", "null"}, []int{-32600, -32600, -32600, -32600}},
} {
t.Run(tc.name, func(t *testing.T) {
a, io := newEthAdapter(), &captureEthIO{}
// A rejected unsubscribe batch must leave this existing subscription live.
if err := a.HandleClientFrame(io, []byte(`{"jsonrpc":"2.0","id":10,"method":"eth_subscribe","params":["newHeads"]}`)); err != nil {
t.Fatal(err)
}
ethAck(t, a, io, numericEthID(t, io.upstream[0]), "0xupstream")
if err := a.HandleClientFrame(io, []byte(`{"jsonrpc":"2.0","id":"pending","method":"eth_chainId"}`)); err != nil {
t.Fatal(err)
}
beforeID := a.idSeq.Load()
io.upstream, io.client = nil, nil
if err := a.HandleClientFrame(io, []byte(tc.body)); err != nil {
t.Fatal(err)
}
if len(io.upstream) != 0 {
t.Fatalf("rejected batch reached upstream: %q", io.upstream)
}
if a.idSeq.Load() != beforeID || len(a.subs) != 1 || len(a.upToSyn) != 1 || len(a.pending) != 0 || len(a.rpcPending) != 1 {
t.Fatal("rejected batch changed correlation or subscription state")
}
if len(tc.ids) == 0 {
if len(io.client) != 0 {
t.Fatalf("notification-only batch received a reply: %q", io.client)
}
} else {
if len(io.client) != 1 {
t.Fatalf("expected one batch reply, got %q", io.client)
}
var replies []struct {
ID json.RawMessage `json:"id"`
JSONRPC string `json:"jsonrpc"`
Error struct {
Code int `json:"code"`
Message string `json:"message"`
} `json:"error"`
}
if err := json.Unmarshal(io.client[0], &replies); err != nil {
t.Fatalf("reply is not an array: %s", io.client[0])
}
if len(replies) != len(tc.ids) {
t.Fatalf("wrong number of replies: %s", io.client[0])
}
for i, reply := range replies {
if string(reply.ID) != tc.ids[i] || reply.JSONRPC != "2.0" || reply.Error.Code != tc.codes[i] || reply.Error.Message == "" {
t.Fatalf("incorrect reply %d: %s", i, io.client[0])
}
}
}
io.client = nil
if err := a.HandleUpstreamFrame(io, []byte(`{"jsonrpc":"2.0","method":"eth_subscription","params":{"subscription":"0xupstream","result":{"number":"0x64"}}}`)); err != nil {
t.Fatal(err)
}
if len(io.client) != 1 {
t.Fatal("existing subscription stopped delivering after batch rejection")
}
})
}
}

func TestEthNonSubscriptionBatchHandlingIsUnchanged(t *testing.T) {
for _, body := range []string{`[]`, `[`, `[1]`, `[{"jsonrpc":"2.0","method":"eth_chainId"}]`} {
t.Run(body, func(t *testing.T) {
a, io := newEthAdapter(), &captureEthIO{}
if err := a.HandleClientFrame(io, []byte(body)); err != nil {
t.Fatal(err)
}
if len(io.client) != 0 || len(io.upstream) != 1 || string(io.upstream[0]) != body {
t.Fatalf("ordinary/invalid batch handling changed: upstream=%q client=%q", io.upstream, io.client)
}
})
}
}
205 changes: 205 additions & 0 deletions internal/subscription/eth_ids_test.go
Original file line number Diff line number Diff line change
@@ -0,0 +1,205 @@
package subscription

import (
"context"
"encoding/json"
"strconv"
"testing"
)

type captureEthIO struct {
upstream [][]byte
client [][]byte
}

func (io *captureEthIO) upstreamWrite(msg []byte) error {
io.upstream = append(io.upstream, append([]byte(nil), msg...))
return nil
}

func (io *captureEthIO) clientWrite(msg []byte) error {
io.client = append(io.client, append([]byte(nil), msg...))
return nil
}

func (io *captureEthIO) reply(id json.RawMessage, key string, value any) error {
msg, err := json.Marshal(map[string]any{"jsonrpc": "2.0", "id": id, key: value})
if err != nil {
return err
}
return io.clientWrite(msg)
}

func (io *captureEthIO) clientReplyResult(id json.RawMessage, result string) error {
return io.reply(id, "result", result)
}

func (io *captureEthIO) clientReplyBool(id json.RawMessage, result bool) error {
return io.reply(id, "result", result)
}

func (io *captureEthIO) clientReplyError(id json.RawMessage, code int, msg string) error {
return io.reply(id, "error", map[string]any{"code": code, "message": msg})
}

func ethEnvelope(t *testing.T, msg []byte) map[string]json.RawMessage {
t.Helper()
var result map[string]json.RawMessage
if err := json.Unmarshal(msg, &result); err != nil {
t.Fatalf("invalid JSON %s: %v", msg, err)
}
return result
}

func numericEthID(t *testing.T, msg []byte) json.RawMessage {
t.Helper()
id := ethEnvelope(t, msg)["id"]
if _, err := strconv.ParseUint(string(id), 10, 64); err != nil {
t.Fatalf("upstream ID must be an unquoted integer, got %s", id)
}
return id
}

func ethAck(t *testing.T, a *ethAdapter, io *captureEthIO, id json.RawMessage, result string) {
t.Helper()
msg, err := json.Marshal(map[string]any{"jsonrpc": "2.0", "id": id, "result": result})
if err != nil {
t.Fatal(err)
}
if err := a.HandleUpstreamFrame(io, msg); err != nil {
t.Fatal(err)
}
}

func TestEthSubscriptionNumericUpstreamIDsPreserveClientIDs(t *testing.T) {
for _, clientID := range []string{`1`, `"subscription-request"`, `9007199254740993`, `null`} {
t.Run(clientID, func(t *testing.T) {
a, io := newEthAdapter(), &captureEthIO{}
req := []byte(`{"jsonrpc":"2.0","id":` + clientID + `,"method":"eth_subscribe","params":["newHeads"]}`)
if err := a.HandleClientFrame(io, req); err != nil {
t.Fatal(err)
}
firstID := numericEthID(t, io.upstream[0])
ethAck(t, a, io, firstID, "0xupstream")
response := ethEnvelope(t, io.client[0])
if string(response["id"]) != clientID || string(response["result"]) != `"0x0000000000000001"` {
t.Fatalf("unexpected client subscribe response: %s", io.client[0])
}
// A real notification advances the cursor before reconnect/replay.
notification := []byte(`{"jsonrpc":"2.0","method":"eth_subscription","params":{"subscription":"0xupstream","result":{"number":"0x64"}}}`)
if err := a.HandleUpstreamFrame(io, notification); err != nil {
t.Fatal(err)
}
if err := a.ReplaySubs(context.Background(), io); err != nil {
t.Fatal(err)
}
replayID := numericEthID(t, io.upstream[1])
if string(firstID) == string(replayID) {
t.Fatal("replay reused the original request ID")
}
ethAck(t, a, io, replayID, "0xresumed")
if len(io.client) != 2 {
t.Fatalf("replay generated an extra client reply: %q", io.client)
}
unsubscribe := []byte(`{"jsonrpc":"2.0","id":` + clientID + `,"method":"eth_unsubscribe","params":["0x0000000000000001"]}`)
if err := a.HandleClientFrame(io, unsubscribe); err != nil {
t.Fatal(err)
}
unsubscribeID := numericEthID(t, io.upstream[2])
response = ethEnvelope(t, io.client[2])
if string(response["id"]) != clientID || string(response["result"]) != "true" {
t.Fatalf("unexpected unsubscribe reply: %s", io.client[2])
}
if err := a.HandleUpstreamFrame(io, []byte(`{"jsonrpc":"2.0","id":`+string(unsubscribeID)+`,"result":true}`)); err != nil {
t.Fatal(err)
}
if len(io.client) != 3 {
t.Fatalf("internal unsubscribe ack leaked to client: %q", io.client)
}
})
}
}

func TestEthOrdinaryCallCannotCollideWithPendingSubscribe(t *testing.T) {
a, io := newEthAdapter(), &captureEthIO{}
for _, req := range []string{
`{"jsonrpc":"2.0","id":1,"method":"eth_subscribe","params":["newHeads"]}`,
`{"jsonrpc":"2.0","id":1,"method":"eth_chainId","params":[]}`,
} {
if err := a.HandleClientFrame(io, []byte(req)); err != nil {
t.Fatal(err)
}
}
subID := numericEthID(t, io.upstream[0])
callID := numericEthID(t, io.upstream[1])
if string(subID) == string(callID) {
t.Fatal("ordinary call collided with pending subscribe")
}
// The ordinary reply arrives before the subscribe acknowledgement.
ethAck(t, a, io, callID, "0x59f")
ethAck(t, a, io, subID, "0xsubscription")
for index, wantResult := range []string{`"0x59f"`, `"0x0000000000000001"`} {
reply := ethEnvelope(t, io.client[index])
if string(reply["id"]) != "1" || string(reply["result"]) != wantResult {
t.Fatalf("reply %d was miscorrelated: %s", index, io.client[index])
}
}
}

func TestEthBatchIDsPreserveTypePrecisionAndResponseShape(t *testing.T) {
a, io := newEthAdapter(), &captureEthIO{}
clientIDs := []string{`1`, `"1"`, `9007199254740993`, `null`}
batch := make([]json.RawMessage, 0, len(clientIDs))
for _, id := range clientIDs {
batch = append(batch, json.RawMessage(`{"jsonrpc":"2.0","id":`+id+`,"method":"eth_chainId","params":[]}`))
}
req, err := json.Marshal(batch)
if err != nil {
t.Fatal(err)
}
if err := a.HandleClientFrame(io, req); err != nil {
t.Fatal(err)
}
var upstream []json.RawMessage
if err := json.Unmarshal(io.upstream[0], &upstream); err != nil {
t.Fatal(err)
}
responses := make([]json.RawMessage, 0, len(upstream))
for i := len(upstream) - 1; i >= 0; i-- {
id := numericEthID(t, upstream[i])
responses = append(responses, json.RawMessage(`{"jsonrpc":"2.0","id":`+string(id)+`,"result":"0x59f"}`))
}
response, err := json.Marshal(responses)
if err != nil {
t.Fatal(err)
}
if err := a.HandleUpstreamFrame(io, response); err != nil {
t.Fatal(err)
}
var restored []json.RawMessage
if err := json.Unmarshal(io.client[0], &restored); err != nil {
t.Fatalf("batch response shape was lost: %v", err)
}
for i, msg := range restored {
id := ethEnvelope(t, msg)["id"]
if string(id) != clientIDs[len(clientIDs)-1-i] {
t.Fatalf("response ID lost type/precision: %s", id)
}
}
}

func TestEthReplayDiscardsAbandonedOrdinaryCallIDs(t *testing.T) {
a, io := newEthAdapter(), &captureEthIO{}
if err := a.HandleClientFrame(io, []byte(`{"jsonrpc":"2.0","id":"unfinished","method":"eth_chainId","params":[]}`)); err != nil {
t.Fatal(err)
}
if len(a.rpcPending) != 1 {
t.Fatal("ordinary request did not register correlation")
}
if err := a.ReplaySubs(context.Background(), io); err != nil {
t.Fatal(err)
}
if len(a.rpcPending) != 0 {
t.Fatal("dead upstream correlation entries survived reconnect")
}
}
Loading
Loading