perf: sync protocol and core hardening updates
This commit is contained in:
parent
152fed3b87
commit
4390ebf5a9
283 changed files with 29231 additions and 2295 deletions
|
|
@ -20,6 +20,7 @@
|
|||
| `schema/canonical-227.tl` | **embed**,运行期 walker 的 227 字段布局(= gotd `td/_schema/tdesktop.tl` 的副本) | gotd 升级时 re-sync |
|
||||
| `_schema/layer-2NN.tl` | 历史层官方 schema(从 TDesktop git 抽,**仅生成期用**,下划线=不编译/不 embed) | 升级/下探 floor 时抽取 |
|
||||
| `schema/client-drift.tl` | **声明式**:客户端发的旧构造器老布局(body 与 227 不同的) | 发现客户端漂移时 +1 行 |
|
||||
| `schema/routable-compat.tl` | **仅结构预检**:已有 RPC fallback adapter 的非 canonical wire 布局(当前只含 4 个 DrKLO theme 构造器);与 canonical 图合并后完整 walk,但不自动升级 | 收敛既有手写 adapter 时维护,禁止借此新增业务 fallback |
|
||||
| `client_aliases.go` | 客户端漂移里 **body 与 227 字节一致**的,纯 `老CRC→227CRC` | 发现纯换 CRC 漂移时 +1 条 |
|
||||
| `tables_gen.go` | **生成产物**(勿手改):官方层降级表 + 入站升级表 + 新类型集 | 跑 `gen` 重生成 |
|
||||
| `gen/main.go` | 生成器:对拍 schema、证明机械性、产 `tables_gen.go` | 升级逻辑变更时 |
|
||||
|
|
@ -95,7 +96,7 @@ gofmt -w internal/compat/layerwire/ && go build ./... && go vet ./internal/...
|
|||
- 绿 = 通用引擎已能自动升级(复制共享字段 + 插 flags=0 + 按 kind 补默认)。**完事**。
|
||||
- `TestInboundDriftCoverage` 报 `needs converter A->B` = 有字段类型变更 → 往 `inbound.go fieldConverters` 加一条 `"A->B"`(可复用,参照 `Vector<int>->Vector<InputMessage>`)。
|
||||
- 报 `field X not defaultable` 或字段**改名** → 往 `inbound.go driftFieldRenames` 加 `"<method>\x00<227字段>": "<老字段>"`(参照 `bots.exportBotToken\x00bot`)。
|
||||
4. **绝不**为此写一个新的 `handleLegacyXxx` 解码 handler——那是旧做法,已全删。统一走数据 + 通用引擎。
|
||||
4. **绝不**为此写一个新的 `handleLegacyXxx` 解码 handler——统一走数据 + 通用引擎。`routable-compat.tl` 只给既存 DrKLO theme fallback 补 dispatcher 前结构门禁,不是新增 adapter 的入口。
|
||||
|
||||
## 操作 4:出站 `TestCoverageGate` 失败
|
||||
|
||||
|
|
|
|||
|
|
@ -30,8 +30,11 @@ func init() {
|
|||
// replaceWithBare consumes the canonical (227-only) object and emits a
|
||||
// bodyless constructor id the target layer understands.
|
||||
func replaceWithBare(id uint32) fallbackFunc {
|
||||
return func(cl *ctorLayout, in, out *bin.Buffer, layer int) error {
|
||||
if err := canonical.skipObject(in); err != nil {
|
||||
return func(cl *ctorLayout, in, out *bin.Buffer, layer, depth int, walk *walkState) error {
|
||||
if err := in.ConsumeID(cl.crc); err != nil {
|
||||
return err
|
||||
}
|
||||
if err := walk.skipCtorBody(canonical, in, cl, depth); err != nil {
|
||||
return err
|
||||
}
|
||||
out.PutID(id)
|
||||
|
|
@ -47,6 +50,8 @@ var peerVectorField = fieldLayout{
|
|||
elem: &fieldLayout{kind: kindObject, typeName: "Peer", flagBit: -1},
|
||||
}
|
||||
|
||||
var pollOptionBytesField = fieldLayout{kind: kindBytes, flagBit: -1}
|
||||
|
||||
// transcodePollAnswerVoters downgrades pollAnswerVoters: canonical (227) made
|
||||
// voters conditional (flags.2?int) and added recent_voters (flags.2?Vector<Peer>);
|
||||
// older layers carry voters as a plain int. The leading CRC is already consumed.
|
||||
|
|
@ -54,34 +59,35 @@ var peerVectorField = fieldLayout{
|
|||
// 227: flags:# chosen:flags.0?true correct:flags.1?true option:bytes
|
||||
// voters:flags.2?int recent_voters:flags.2?Vector<Peer>
|
||||
// <=226: flags:# chosen:flags.0?true correct:flags.1?true option:bytes voters:int
|
||||
func transcodePollAnswerVoters(cl *ctorLayout, target uint32, in, out *bin.Buffer, layer int) error {
|
||||
func transcodePollAnswerVoters(cl *ctorLayout, target uint32, in, out *bin.Buffer, layer, depth int, walk *walkState) error {
|
||||
flags, err := in.Uint32()
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
option, err := in.Bytes()
|
||||
if err != nil {
|
||||
optionStart := in.Buf
|
||||
if err := walk.skipValue(canonical, in, &pollOptionBytesField, cl, depth); err != nil {
|
||||
return err
|
||||
}
|
||||
optionRaw := optionStart[:len(optionStart)-len(in.Buf)]
|
||||
var voters int
|
||||
if flags&(1<<2) != 0 {
|
||||
if voters, err = in.Int(); err != nil {
|
||||
return err
|
||||
}
|
||||
if err := canonical.skipValue(in, &peerVectorField); err != nil {
|
||||
if err := walk.skipValue(canonical, in, &peerVectorField, cl, depth); err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
out.PutID(target)
|
||||
out.PutUint32(flags & 0b11) // retain chosen/correct, clear the moved bit 2
|
||||
out.PutBytes(option)
|
||||
out.Put(optionRaw)
|
||||
out.PutInt(voters)
|
||||
return nil
|
||||
}
|
||||
|
||||
// fallbackMessageEntity replaces any 227-only MessageEntity with
|
||||
// messageEntityUnknown, preserving offset/length so text positions stay valid.
|
||||
func fallbackMessageEntity(cl *ctorLayout, in, out *bin.Buffer, layer int) error {
|
||||
func fallbackMessageEntity(cl *ctorLayout, in, out *bin.Buffer, layer, depth int, walk *walkState) error {
|
||||
id, err := in.PeekID()
|
||||
if err != nil {
|
||||
return err
|
||||
|
|
@ -89,7 +95,7 @@ func fallbackMessageEntity(cl *ctorLayout, in, out *bin.Buffer, layer int) error
|
|||
if err := in.ConsumeID(id); err != nil {
|
||||
return err
|
||||
}
|
||||
offset, length, err := canonical.decodeOffsetLength(in, cl)
|
||||
offset, length, err := canonical.decodeOffsetLength(in, cl, depth, walk)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
|
@ -101,7 +107,7 @@ func fallbackMessageEntity(cl *ctorLayout, in, out *bin.Buffer, layer int) error
|
|||
|
||||
// decodeOffsetLength walks a constructor body (no leading CRC) per the canonical
|
||||
// layout, returning its offset/length int fields and discarding the rest.
|
||||
func (m *schemaModel) decodeOffsetLength(in *bin.Buffer, cl *ctorLayout) (offset, length int, err error) {
|
||||
func (m *schemaModel) decodeOffsetLength(in *bin.Buffer, cl *ctorLayout, depth int, walk *walkState) (offset, length int, err error) {
|
||||
var flags map[string]uint32
|
||||
for i := range cl.fields {
|
||||
f := &cl.fields[i]
|
||||
|
|
@ -129,7 +135,7 @@ func (m *schemaModel) decodeOffsetLength(in *bin.Buffer, cl *ctorLayout) (offset
|
|||
return
|
||||
}
|
||||
default:
|
||||
if err = m.skipValue(in, f); err != nil {
|
||||
if err = walk.skipValue(m, in, f, cl, depth); err != nil {
|
||||
return
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -47,16 +47,19 @@ var driftFieldRenames = map[string]string{
|
|||
// fieldConverter rewrites one field whose wire type changed between the old and
|
||||
// canonical layout. Keyed by "<oldTypeSig>-><newTypeSig>"; raw is the old field's
|
||||
// encoded bytes. Reusable across any method with the same type change.
|
||||
type fieldConverter func(raw []byte, out *bin.Buffer) error
|
||||
type fieldConverter func(raw []byte, out *bin.Buffer, walk *walkState, owner *ctorLayout, field *fieldLayout) error
|
||||
|
||||
var fieldConverters = map[string]fieldConverter{
|
||||
// id:Vector<int> -> id:Vector<InputMessage> (wrap each int in inputMessageID).
|
||||
"Vector<int>->Vector<InputMessage>": func(raw []byte, out *bin.Buffer) error {
|
||||
"Vector<int>->Vector<InputMessage>": func(raw []byte, out *bin.Buffer, walk *walkState, owner *ctorLayout, field *fieldLayout) error {
|
||||
in := &bin.Buffer{Buf: raw}
|
||||
n, err := in.VectorHeader()
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if max := walk.vectorLimit(owner, field); n > max {
|
||||
return limitf("vector %s.%s length %d exceeds limit %d", ownerName(owner), fieldName(field), n, max)
|
||||
}
|
||||
out.PutVectorHeader(n)
|
||||
for i := 0; i < n; i++ {
|
||||
v, err := in.Int()
|
||||
|
|
@ -66,10 +69,13 @@ var fieldConverters = map[string]fieldConverter{
|
|||
out.PutID(inputMessageID)
|
||||
out.PutInt(v)
|
||||
}
|
||||
if in.Len() != 0 {
|
||||
return malformedf("%d trailing bytes in Vector<int> converter", in.Len())
|
||||
}
|
||||
return nil
|
||||
},
|
||||
// bot_id:long -> bot:InputUser{user_id, access_hash=0}.
|
||||
"long->InputUser": func(raw []byte, out *bin.Buffer) error {
|
||||
"long->InputUser": func(raw []byte, out *bin.Buffer, walk *walkState, owner *ctorLayout, field *fieldLayout) error {
|
||||
in := &bin.Buffer{Buf: raw}
|
||||
id, err := in.Long()
|
||||
if err != nil {
|
||||
|
|
@ -78,11 +84,14 @@ var fieldConverters = map[string]fieldConverter{
|
|||
out.PutID(inputUserID)
|
||||
out.PutLong(id)
|
||||
out.PutLong(0)
|
||||
if in.Len() != 0 {
|
||||
return malformedf("%d trailing bytes in long converter", in.Len())
|
||||
}
|
||||
return nil
|
||||
},
|
||||
// channel:InputChannel -> peer:InputPeer for the old channels.editCreator
|
||||
// Android constructor. Concrete layouts are otherwise byte-compatible.
|
||||
"InputChannel->InputPeer": func(raw []byte, out *bin.Buffer) error {
|
||||
"InputChannel->InputPeer": func(raw []byte, out *bin.Buffer, walk *walkState, owner *ctorLayout, field *fieldLayout) error {
|
||||
in := &bin.Buffer{Buf: raw}
|
||||
id, err := in.ID()
|
||||
if err != nil {
|
||||
|
|
@ -116,8 +125,12 @@ var fieldConverters = map[string]fieldConverter{
|
|||
// id + body) is what to dispatch.
|
||||
func UpgradeInbound(id uint32, in *bin.Buffer) (*bin.Buffer, bool, error) {
|
||||
if newID, ok := UpgradeMethodCRC(id); ok {
|
||||
if len(in.Buf) < 4 {
|
||||
return nil, false, fmt.Errorf("layerwire: short inbound buffer for %#08x", id)
|
||||
target := canonical.byCRC[newID]
|
||||
if target == nil || !target.isFunc {
|
||||
return nil, true, malformedf("alias %#08x targets unknown canonical method %#08x", id, newID)
|
||||
}
|
||||
if err := validateAliasedMethod(id, target, in.Buf); err != nil {
|
||||
return nil, true, err
|
||||
}
|
||||
// Copy rather than rewrite in place: never mutate the caller's buffer
|
||||
// (matches the body-transform path, which also returns a fresh buffer).
|
||||
|
|
@ -126,15 +139,36 @@ func UpgradeInbound(id uint32, in *bin.Buffer) (*bin.Buffer, bool, error) {
|
|||
return out, true, nil
|
||||
}
|
||||
if old := driftModel.byCRC[id]; old != nil {
|
||||
out, err := upgradeFromDrift(old, in)
|
||||
out, err := upgradeFromDrift(old, in, newWalkState())
|
||||
if err != nil {
|
||||
return nil, false, fmt.Errorf("layerwire: upgrade %s (%#08x): %w", old.name, id, err)
|
||||
return nil, true, classifyWalkError(fmt.Errorf("layerwire: upgrade %s (%#08x): %w", old.name, id, err))
|
||||
}
|
||||
return out, true, nil
|
||||
}
|
||||
return nil, false, nil
|
||||
}
|
||||
|
||||
// validateAliasedMethod validates the old-id/canonical-body shape before
|
||||
// allocating the replacement buffer. The body is walked against the canonical
|
||||
// target layout while the original constructor id remains untouched.
|
||||
func validateAliasedMethod(oldID uint32, target *ctorLayout, raw []byte) error {
|
||||
walk := newWalkState()
|
||||
if err := walk.enter(1, "constructor"); err != nil {
|
||||
return err
|
||||
}
|
||||
b := &bin.Buffer{Buf: raw}
|
||||
if err := b.ConsumeID(oldID); err != nil {
|
||||
return classifyWalkError(err)
|
||||
}
|
||||
if err := walk.skipCtorBody(canonical, b, target, 1); err != nil {
|
||||
return classifyWalkError(err)
|
||||
}
|
||||
if b.Len() != 0 {
|
||||
return malformedf("%d trailing bytes after aliased method %s", b.Len(), target.name)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// IsClientDrift reports whether id is a client-private constructor (DrKLO
|
||||
// constructor drift), as opposed to official layer drift from api.tl.
|
||||
func IsClientDrift(id uint32) bool {
|
||||
|
|
@ -146,11 +180,14 @@ func IsClientDrift(id uint32) bool {
|
|||
|
||||
// upgradeFromDrift rebuilds a canonical (227) request from an old client-drift
|
||||
// body, comparing the declared old layout to the canonical layout field by field.
|
||||
func upgradeFromDrift(old *ctorLayout, in *bin.Buffer) (*bin.Buffer, error) {
|
||||
func upgradeFromDrift(old *ctorLayout, in *bin.Buffer, walk *walkState) (*bin.Buffer, error) {
|
||||
target := canonical.byName[old.name]
|
||||
if target == nil {
|
||||
return nil, fmt.Errorf("no canonical method %q", old.name)
|
||||
}
|
||||
if err := walk.enter(1, "constructor"); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if err := in.ConsumeID(old.crc); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
|
@ -179,7 +216,7 @@ func upgradeFromDrift(old *ctorLayout, in *bin.Buffer) (*bin.Buffer, error) {
|
|||
continue
|
||||
}
|
||||
pre := in.Buf
|
||||
if err := canonical.skipValue(in, f); err != nil {
|
||||
if err := walk.skipValue(canonical, in, f, old, 1); err != nil {
|
||||
return nil, fmt.Errorf("decode old field %q: %w", f.name, err)
|
||||
}
|
||||
vals[f.name] = pre[:len(pre)-len(in.Buf)]
|
||||
|
|
@ -208,7 +245,7 @@ func upgradeFromDrift(old *ctorLayout, in *bin.Buffer) (*bin.Buffer, error) {
|
|||
if conv == nil {
|
||||
return nil, fmt.Errorf("field %q: no converter %s->%s", nf.name, typeSig(of), typeSig(nf))
|
||||
}
|
||||
if err := conv(vals[oldName], out); err != nil {
|
||||
if err := conv(vals[oldName], out, walk, old, of); err != nil {
|
||||
return nil, fmt.Errorf("field %q convert: %w", nf.name, err)
|
||||
}
|
||||
} else {
|
||||
|
|
|
|||
80
internal/compat/layerwire/routable.go
Normal file
80
internal/compat/layerwire/routable.go
Normal file
|
|
@ -0,0 +1,80 @@
|
|||
package layerwire
|
||||
|
||||
import (
|
||||
_ "embed"
|
||||
"fmt"
|
||||
|
||||
"github.com/gotd/td/bin"
|
||||
)
|
||||
|
||||
const maxOpaqueRequestBytes = 16 << 20
|
||||
|
||||
//go:embed schema/routable-compat.tl
|
||||
var routableCompatSchema string
|
||||
|
||||
// routable combines the canonical Layer 227 model with the small set of
|
||||
// explicitly declared compatibility-only methods. Nested objects in those
|
||||
// methods are canonical Input* constructors, so one combined graph is needed
|
||||
// for the same depth/vector/bytes walker to validate the complete request.
|
||||
var routable = mustLoadRoutable()
|
||||
|
||||
func mustLoadRoutable() *schemaModel {
|
||||
compat, err := parseSchemaModel(routableCompatSchema)
|
||||
if err != nil {
|
||||
panic("layerwire: parse routable compat schema: " + err.Error())
|
||||
}
|
||||
m := &schemaModel{
|
||||
byCRC: make(map[uint32]*ctorLayout, len(canonical.byCRC)+len(compat.byCRC)),
|
||||
byName: make(map[string]*ctorLayout, len(canonical.byName)+len(compat.byName)),
|
||||
bareByT: make(map[string]*ctorLayout, len(canonical.bareByT)),
|
||||
ctorsOfT: make(map[string][]*ctorLayout, len(canonical.ctorsOfT)),
|
||||
}
|
||||
for id, cl := range canonical.byCRC {
|
||||
m.byCRC[id] = cl
|
||||
}
|
||||
for name, cl := range canonical.byName {
|
||||
m.byName[name] = cl
|
||||
}
|
||||
for name, cl := range canonical.bareByT {
|
||||
m.bareByT[name] = cl
|
||||
}
|
||||
for name, ctors := range canonical.ctorsOfT {
|
||||
m.ctorsOfT[name] = ctors
|
||||
}
|
||||
for id, cl := range compat.byCRC {
|
||||
if existing := m.byCRC[id]; existing != nil {
|
||||
panic(fmt.Sprintf("layerwire: routable compat crc %#08x collides with %s", id, existing.name))
|
||||
}
|
||||
m.byCRC[id] = cl
|
||||
m.byName[cl.name] = cl
|
||||
}
|
||||
return m
|
||||
}
|
||||
|
||||
// ValidateRoutableRequest validates every request shape the router knows how to
|
||||
// decode, including compatibility-only fallback methods. known=false denotes
|
||||
// a genuinely unknown top-level constructor. Such a request is never decoded:
|
||||
// it is treated as opaque, word-aligned TL data, bounded by both this total-size
|
||||
// cap and mtprotoedge's transport/RPC budgets, and must continue to the router's
|
||||
// compatibility trace rather than being mislabeled as malformed input.
|
||||
func ValidateRoutableRequest(body []byte) (known bool, err error) {
|
||||
b := &bin.Buffer{Buf: body}
|
||||
id, err := b.PeekID()
|
||||
if err != nil {
|
||||
return false, classifyWalkError(err)
|
||||
}
|
||||
cl := routable.byCRC[id]
|
||||
if cl == nil {
|
||||
if len(body) > maxOpaqueRequestBytes {
|
||||
return false, limitf("opaque request length %d exceeds limit %d", len(body), maxOpaqueRequestBytes)
|
||||
}
|
||||
if len(body)%bin.Word != 0 {
|
||||
return false, malformedf("opaque request length %d is not word aligned", len(body))
|
||||
}
|
||||
return false, nil
|
||||
}
|
||||
if !cl.isFunc {
|
||||
return true, malformedf("constructor %s (%#08x) is not a method", cl.name, id)
|
||||
}
|
||||
return true, validateRequestLayout(routable, cl, body)
|
||||
}
|
||||
43
internal/compat/layerwire/routable_test.go
Normal file
43
internal/compat/layerwire/routable_test.go
Normal file
|
|
@ -0,0 +1,43 @@
|
|||
package layerwire
|
||||
|
||||
import (
|
||||
"errors"
|
||||
"testing"
|
||||
|
||||
"github.com/gotd/td/bin"
|
||||
"github.com/gotd/td/tg"
|
||||
)
|
||||
|
||||
func TestValidateRoutableRequestCompatibilityAndUnknown(t *testing.T) {
|
||||
t.Run("legacy theme is fully walked", func(t *testing.T) {
|
||||
var b bin.Buffer
|
||||
b.PutID(0x8d9d742b)
|
||||
b.PutString("android")
|
||||
(&tg.InputThemeSlug{Slug: "night"}).Encode(&b)
|
||||
b.PutLong(42)
|
||||
known, err := ValidateRoutableRequest(b.Buf)
|
||||
if err != nil || !known {
|
||||
t.Fatalf("legacy theme known=%v err=%v, want true/nil", known, err)
|
||||
}
|
||||
|
||||
b.Buf = b.Buf[:len(b.Buf)-4]
|
||||
known, err = ValidateRoutableRequest(b.Buf)
|
||||
if !known || !errors.Is(err, ErrMalformed) {
|
||||
t.Fatalf("truncated legacy theme known=%v err=%v, want true/malformed", known, err)
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("unknown stays opaque and bounded", func(t *testing.T) {
|
||||
var b bin.Buffer
|
||||
b.PutID(0x12345678)
|
||||
b.PutUint32(0xffffffff)
|
||||
known, err := ValidateRoutableRequest(b.Buf)
|
||||
if err != nil || known {
|
||||
t.Fatalf("opaque unknown known=%v err=%v, want false/nil", known, err)
|
||||
}
|
||||
known, err = ValidateRoutableRequest(append(b.Buf, 1))
|
||||
if known || !errors.Is(err, ErrMalformed) {
|
||||
t.Fatalf("unaligned unknown known=%v err=%v, want false/malformed", known, err)
|
||||
}
|
||||
})
|
||||
}
|
||||
11
internal/compat/layerwire/schema/routable-compat.tl
Normal file
11
internal/compat/layerwire/schema/routable-compat.tl
Normal file
|
|
@ -0,0 +1,11 @@
|
|||
// Hand-maintained request layouts that are intentionally handled by the RPC
|
||||
// fallback instead of gotd's canonical ServerDispatcher. They still belong in
|
||||
// the structural preflight model: fallback handlers must never become a way to
|
||||
// bypass the canonical vector/depth/bytes budgets.
|
||||
|
||||
---functions---
|
||||
|
||||
compat.legacyCreateTheme#8432c21f flags:# slug:string title:string document:flags.2?InputDocument settings:flags.3?InputThemeSettings = Object;
|
||||
compat.legacyUpdateTheme#5cb367d5 flags:# format:string theme:InputTheme slug:flags.0?string title:flags.1?string document:flags.2?InputDocument settings:flags.3?InputThemeSettings = Object;
|
||||
compat.legacyInstallTheme#7ae43737 flags:# dark:flags.0?true format:flags.1?string theme:flags.1?InputTheme = Object;
|
||||
compat.legacyGetTheme#8d9d742b format:string theme:InputTheme document_id:long = Object;
|
||||
|
|
@ -139,11 +139,11 @@ func (lt *layerTables) fieldDirty(f *fieldLayout) bool {
|
|||
// field drop. The leading CRC has already been consumed from in; the transform
|
||||
// reads the canonical body from in and writes the target-layer object (whose
|
||||
// constructor id is target) to out.
|
||||
type structuralFunc func(cl *ctorLayout, target uint32, in, out *bin.Buffer, layer int) error
|
||||
type structuralFunc func(cl *ctorLayout, target uint32, in, out *bin.Buffer, layer, depth int, walk *walkState) error
|
||||
|
||||
// fallbackFunc replaces a layer-absent (227-only) constructor with an
|
||||
// equivalent the target layer understands. The leading CRC is NOT yet consumed.
|
||||
type fallbackFunc func(cl *ctorLayout, in, out *bin.Buffer, layer int) error
|
||||
type fallbackFunc func(cl *ctorLayout, in, out *bin.Buffer, layer, depth int, walk *walkState) error
|
||||
|
||||
// structuralTransforms and the newType fallback registries are populated in
|
||||
// fallback.go. newTypeFallbacks is keyed by canonical CRC (specific override);
|
||||
|
|
@ -175,11 +175,12 @@ func Transcode(canonicalBytes []byte, layer int) ([]byte, error) {
|
|||
}
|
||||
in := &bin.Buffer{Buf: canonicalBytes}
|
||||
out := &bin.Buffer{}
|
||||
if err := lt.transcodeObject(in, out, layer); err != nil {
|
||||
return nil, err
|
||||
walk := newWalkState()
|
||||
if err := lt.transcodeObject(in, out, layer, 1, walk); err != nil {
|
||||
return nil, classifyWalkError(err)
|
||||
}
|
||||
if in.Len() != 0 {
|
||||
return nil, fmt.Errorf("layerwire: %d trailing bytes after transcode to layer %d", in.Len(), layer)
|
||||
return nil, malformedf("%d trailing bytes after transcode to layer %d", in.Len(), layer)
|
||||
}
|
||||
return out.Buf, nil
|
||||
}
|
||||
|
|
@ -198,7 +199,10 @@ func UpgradeMethodCRC(oldID uint32) (uint32, bool) {
|
|||
return newID, ok
|
||||
}
|
||||
|
||||
func (lt *layerTables) transcodeObject(in, out *bin.Buffer, layer int) error {
|
||||
func (lt *layerTables) transcodeObject(in, out *bin.Buffer, layer, depth int, walk *walkState) error {
|
||||
if err := walk.enter(depth, "constructor"); err != nil {
|
||||
return err
|
||||
}
|
||||
id, err := in.PeekID()
|
||||
if err != nil {
|
||||
return err
|
||||
|
|
@ -216,10 +220,10 @@ func (lt *layerTables) transcodeObject(in, out *bin.Buffer, layer int) error {
|
|||
if fn == nil {
|
||||
return fmt.Errorf("layerwire: no structural transform %q for %s@%d", rule.structural, cl.name, layer)
|
||||
}
|
||||
return fn(cl, rule.target, in, out, layer)
|
||||
return fn(cl, rule.target, in, out, layer, depth, walk)
|
||||
}
|
||||
out.PutID(rule.target)
|
||||
return lt.transcodeBody(in, out, cl, rule.keep, layer)
|
||||
return lt.transcodeBody(in, out, cl, rule.keep, layer, depth, walk)
|
||||
}
|
||||
if lt.newTypes[id] {
|
||||
fn := newTypeFallbacks[id]
|
||||
|
|
@ -229,12 +233,15 @@ func (lt *layerTables) transcodeObject(in, out *bin.Buffer, layer int) error {
|
|||
if fn == nil {
|
||||
return fmt.Errorf("layerwire: %s (%#08x) absent at layer %d and no fallback", cl.name, id, layer)
|
||||
}
|
||||
return fn(cl, in, out, layer)
|
||||
return fn(cl, in, out, layer, depth, walk)
|
||||
}
|
||||
if !lt.dirty[id] {
|
||||
// Unaffected subtree: byte-for-byte copy.
|
||||
pre := in.Buf
|
||||
if err := canonical.skipObject(in); err != nil {
|
||||
if err := in.ConsumeID(id); err != nil {
|
||||
return err
|
||||
}
|
||||
if err := walk.skipCtorBody(canonical, in, cl, depth); err != nil {
|
||||
return err
|
||||
}
|
||||
out.Put(pre[:len(pre)-len(in.Buf)])
|
||||
|
|
@ -245,14 +252,14 @@ func (lt *layerTables) transcodeObject(in, out *bin.Buffer, layer int) error {
|
|||
return err
|
||||
}
|
||||
out.PutID(id)
|
||||
return lt.transcodeBody(in, out, cl, nil, layer)
|
||||
return lt.transcodeBody(in, out, cl, nil, layer, depth, walk)
|
||||
}
|
||||
|
||||
// transcodeBody re-encodes a constructor body. keep==nil means retain every
|
||||
// field (recursing into dirty descendants); otherwise only the named canonical
|
||||
// fields are written, flag integers are remasked to the retained bits, and
|
||||
// dropped fields are read-and-discarded.
|
||||
func (lt *layerTables) transcodeBody(in, out *bin.Buffer, cl *ctorLayout, keep map[string]bool, layer int) error {
|
||||
func (lt *layerTables) transcodeBody(in, out *bin.Buffer, cl *ctorLayout, keep map[string]bool, layer, depth int, walk *walkState) error {
|
||||
kept := func(name string) bool { return keep == nil || keep[name] }
|
||||
var flags map[string]uint32
|
||||
for i := range cl.fields {
|
||||
|
|
@ -276,10 +283,10 @@ func (lt *layerTables) transcodeBody(in, out *bin.Buffer, cl *ctorLayout, keep m
|
|||
continue
|
||||
}
|
||||
if kept(f.name) {
|
||||
if err := lt.transcodeValue(in, out, f, layer); err != nil {
|
||||
if err := lt.transcodeValue(in, out, f, cl, layer, depth, walk); err != nil {
|
||||
return fmt.Errorf("%s.%s: %w", cl.name, f.name, err)
|
||||
}
|
||||
} else if err := canonical.skipValue(in, f); err != nil {
|
||||
} else if err := walk.skipValue(canonical, in, f, cl, depth); err != nil {
|
||||
return fmt.Errorf("%s.%s (drop): %w", cl.name, f.name, err)
|
||||
}
|
||||
}
|
||||
|
|
@ -301,10 +308,10 @@ func (lt *layerTables) keptMask(cl *ctorLayout, flagName string, kept func(strin
|
|||
|
||||
// transcodeValue writes one present field value, recursing only into dirty
|
||||
// subtrees and byte-copying everything else.
|
||||
func (lt *layerTables) transcodeValue(in, out *bin.Buffer, f *fieldLayout, layer int) error {
|
||||
func (lt *layerTables) transcodeValue(in, out *bin.Buffer, f *fieldLayout, owner *ctorLayout, layer, depth int, walk *walkState) error {
|
||||
if !lt.fieldDirty(f) {
|
||||
pre := in.Buf
|
||||
if err := canonical.skipValue(in, f); err != nil {
|
||||
if err := walk.skipValue(canonical, in, f, owner, depth); err != nil {
|
||||
return err
|
||||
}
|
||||
out.Put(pre[:len(pre)-len(in.Buf)])
|
||||
|
|
@ -312,6 +319,10 @@ func (lt *layerTables) transcodeValue(in, out *bin.Buffer, f *fieldLayout, layer
|
|||
}
|
||||
switch f.kind {
|
||||
case kindVector, kindVectorBare:
|
||||
vectorDepth := depth + 1
|
||||
if vectorDepth <= 0 || vectorDepth > walk.limits.maxDepth {
|
||||
return limitf("vector nesting depth %d exceeds limit %d", vectorDepth, walk.limits.maxDepth)
|
||||
}
|
||||
if f.kind == kindVector {
|
||||
id, err := in.Uint32()
|
||||
if err != nil {
|
||||
|
|
@ -326,23 +337,36 @@ func (lt *layerTables) transcodeValue(in, out *bin.Buffer, f *fieldLayout, layer
|
|||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if n < 0 {
|
||||
return malformedf("negative vector length %d", n)
|
||||
}
|
||||
if max := walk.vectorLimit(owner, f); n > max {
|
||||
return limitf("vector %s.%s length %d exceeds limit %d", ownerName(owner), fieldName(f), n, max)
|
||||
}
|
||||
if err := walk.addUnits(n, "vector "+ownerName(owner)+"."+fieldName(f)); err != nil {
|
||||
return err
|
||||
}
|
||||
out.PutInt(n)
|
||||
for i := 0; i < n; i++ {
|
||||
if err := lt.transcodeValue(in, out, f.elem, layer); err != nil {
|
||||
if err := lt.transcodeValue(in, out, f.elem, nil, layer, vectorDepth, walk); err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
return nil
|
||||
case kindObject:
|
||||
return lt.transcodeObject(in, out, layer)
|
||||
return lt.transcodeObject(in, out, layer, depth+1, walk)
|
||||
case kindBareObject:
|
||||
bareDepth := depth + 1
|
||||
if err := walk.enter(bareDepth, "bare constructor"); err != nil {
|
||||
return err
|
||||
}
|
||||
cl, ok := canonical.bareByT[f.typeName]
|
||||
if !ok {
|
||||
return fmt.Errorf("unknown bare type %q", f.typeName)
|
||||
}
|
||||
// Bare objects have no CRC and (within 220..227) no changed bare ctor;
|
||||
// recurse all-kept to reach any dirty descendants.
|
||||
return lt.transcodeBody(in, out, cl, nil, layer)
|
||||
return lt.transcodeBody(in, out, cl, nil, layer, bareDepth, walk)
|
||||
default:
|
||||
// Primitive marked dirty should be impossible.
|
||||
return fmt.Errorf("unexpected dirty primitive kind %d", f.kind)
|
||||
|
|
|
|||
|
|
@ -1,32 +1,205 @@
|
|||
package layerwire
|
||||
|
||||
import (
|
||||
"errors"
|
||||
"fmt"
|
||||
"io"
|
||||
"math"
|
||||
|
||||
"github.com/gotd/td/bin"
|
||||
)
|
||||
|
||||
// ErrMalformed identifies invalid or truncated TL wire data. Callers may use
|
||||
// errors.Is to distinguish it from an otherwise well-formed request which was
|
||||
// rejected by a walker resource limit.
|
||||
var ErrMalformed = errors.New("layerwire: malformed TL")
|
||||
|
||||
// ErrResourceLimit identifies structurally valid-looking TL input which would
|
||||
// exceed a walker resource budget.
|
||||
var ErrResourceLimit = errors.New("layerwire: resource limit")
|
||||
|
||||
const (
|
||||
defaultMaxVectorElements = 4096
|
||||
defaultMaxWalkDepth = 32
|
||||
defaultMaxWalkUnits = 131072 // constructors + declared vector elements
|
||||
defaultMaxFieldBytes = 16 << 20
|
||||
defaultMaxTotalBytes = 32 << 20
|
||||
)
|
||||
|
||||
// A very small number of API methods have a documented limit above the
|
||||
// package-wide default. Keeping overrides keyed by constructor and field makes
|
||||
// every exception explicit and prevents a large vector in an unrelated method
|
||||
// from inheriting the larger allowance.
|
||||
type vectorLimitKey struct {
|
||||
owner string
|
||||
field string
|
||||
}
|
||||
|
||||
var vectorElementLimitOverrides = map[vectorLimitKey]int{
|
||||
{owner: "contacts.editCloseFriends", field: "id"}: 5000,
|
||||
{owner: "contacts.setBlocked", field: "id"}: 5000,
|
||||
}
|
||||
|
||||
type walkLimits struct {
|
||||
maxVectorElements int
|
||||
maxDepth int
|
||||
maxUnits uint64
|
||||
maxFieldBytes uint64
|
||||
maxTotalBytes uint64
|
||||
}
|
||||
|
||||
var defaultWalkLimits = walkLimits{
|
||||
maxVectorElements: defaultMaxVectorElements,
|
||||
maxDepth: defaultMaxWalkDepth,
|
||||
maxUnits: defaultMaxWalkUnits,
|
||||
maxFieldBytes: defaultMaxFieldBytes,
|
||||
maxTotalBytes: defaultMaxTotalBytes,
|
||||
}
|
||||
|
||||
// walkState is deliberately request-scoped. Every branch of one transform
|
||||
// shares it, so splitting a large value across nested constructors or vectors
|
||||
// cannot reset the aggregate budgets.
|
||||
type walkState struct {
|
||||
limits walkLimits
|
||||
units uint64
|
||||
bytes uint64
|
||||
}
|
||||
|
||||
func newWalkState() *walkState {
|
||||
return &walkState{limits: defaultWalkLimits}
|
||||
}
|
||||
|
||||
func malformedf(format string, args ...any) error {
|
||||
return fmt.Errorf("%w: %s", ErrMalformed, fmt.Sprintf(format, args...))
|
||||
}
|
||||
|
||||
func limitf(format string, args ...any) error {
|
||||
return fmt.Errorf("%w: %s", ErrResourceLimit, fmt.Sprintf(format, args...))
|
||||
}
|
||||
|
||||
// classifyWalkError makes all public walker/transform failures classifiable,
|
||||
// including errors returned by the low-level gotd bin decoder.
|
||||
func classifyWalkError(err error) error {
|
||||
if err == nil || errors.Is(err, ErrMalformed) || errors.Is(err, ErrResourceLimit) {
|
||||
return err
|
||||
}
|
||||
return fmt.Errorf("%w: %v", ErrMalformed, err)
|
||||
}
|
||||
|
||||
func (s *walkState) enter(depth int, what string) error {
|
||||
if depth <= 0 || depth > s.limits.maxDepth {
|
||||
return limitf("%s nesting depth %d exceeds limit %d", what, depth, s.limits.maxDepth)
|
||||
}
|
||||
return s.addUnits(1, what)
|
||||
}
|
||||
|
||||
func (s *walkState) addUnits(n int, what string) error {
|
||||
if n < 0 {
|
||||
return malformedf("negative %s count %d", what, n)
|
||||
}
|
||||
u := uint64(n)
|
||||
// Subtraction form avoids overflow even if limits are changed later.
|
||||
if s.units > s.limits.maxUnits || u > s.limits.maxUnits-s.units {
|
||||
return limitf("constructor/vector element budget exceeds %d at %s", s.limits.maxUnits, what)
|
||||
}
|
||||
s.units += u
|
||||
return nil
|
||||
}
|
||||
|
||||
func (s *walkState) addBytes(n uint64, what string) error {
|
||||
if n > s.limits.maxFieldBytes {
|
||||
return limitf("%s payload length %d exceeds per-field limit %d", what, n, s.limits.maxFieldBytes)
|
||||
}
|
||||
if s.bytes > s.limits.maxTotalBytes || n > s.limits.maxTotalBytes-s.bytes {
|
||||
return limitf("string/bytes payload budget exceeds %d at %s", s.limits.maxTotalBytes, what)
|
||||
}
|
||||
s.bytes += n
|
||||
return nil
|
||||
}
|
||||
|
||||
func (s *walkState) vectorLimit(owner *ctorLayout, f *fieldLayout) int {
|
||||
if owner != nil && f != nil {
|
||||
if n := vectorElementLimitOverrides[vectorLimitKey{owner: owner.name, field: f.name}]; n > 0 {
|
||||
return n
|
||||
}
|
||||
}
|
||||
return s.limits.maxVectorElements
|
||||
}
|
||||
|
||||
const maxConstructorFlagWords = 8
|
||||
|
||||
type constructorFlagWord struct {
|
||||
name string
|
||||
value uint32
|
||||
}
|
||||
|
||||
// ValidateCanonicalRequest performs a complete, allocation-free structural
|
||||
// preflight of one canonical Layer 227 method request. It is intended for the
|
||||
// router seam immediately before typed dispatch. A successful result means the
|
||||
// walker consumed exactly one known function constructor and all of its body.
|
||||
func ValidateCanonicalRequest(body []byte) error {
|
||||
b := &bin.Buffer{Buf: body}
|
||||
id, err := b.PeekID()
|
||||
if err != nil {
|
||||
return classifyWalkError(err)
|
||||
}
|
||||
cl := canonical.byCRC[id]
|
||||
if cl == nil {
|
||||
return malformedf("unknown canonical request constructor %#08x", id)
|
||||
}
|
||||
if !cl.isFunc {
|
||||
return malformedf("constructor %s (%#08x) is not a method", cl.name, id)
|
||||
}
|
||||
return validateRequestLayout(canonical, cl, body)
|
||||
}
|
||||
|
||||
func validateRequestLayout(m *schemaModel, cl *ctorLayout, body []byte) error {
|
||||
b := &bin.Buffer{Buf: body}
|
||||
s := newWalkState()
|
||||
if err := s.skipObject(m, b, 1); err != nil {
|
||||
return classifyWalkError(err)
|
||||
}
|
||||
if b.Len() != 0 {
|
||||
return malformedf("%d trailing bytes after canonical request %s", b.Len(), cl.name)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// skipObject advances b past one boxed object (CRC + body), resolving the
|
||||
// constructor from the canonical schema.
|
||||
// constructor from m. This compatibility wrapper creates a fresh budget; all
|
||||
// production transforms call the stateful variant directly.
|
||||
func (m *schemaModel) skipObject(b *bin.Buffer) error {
|
||||
return classifyWalkError(newWalkState().skipObject(m, b, 1))
|
||||
}
|
||||
|
||||
func (s *walkState) skipObject(m *schemaModel, b *bin.Buffer, depth int) error {
|
||||
if err := s.enter(depth, "constructor"); err != nil {
|
||||
return err
|
||||
}
|
||||
id, err := b.PeekID()
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
cl, ok := m.byCRC[id]
|
||||
if !ok {
|
||||
return fmt.Errorf("layerwire: unknown constructor %#08x", id)
|
||||
return malformedf("unknown constructor %#08x", id)
|
||||
}
|
||||
if err := b.ConsumeID(id); err != nil {
|
||||
return err
|
||||
}
|
||||
return m.skipCtorBody(b, cl)
|
||||
return s.skipCtorBody(m, b, cl, depth)
|
||||
}
|
||||
|
||||
// skipCtorBody advances b past a constructor body (no leading CRC), evaluating
|
||||
// flag integers so conditional fields are read iff present.
|
||||
func (m *schemaModel) skipCtorBody(b *bin.Buffer, cl *ctorLayout) error {
|
||||
var flags map[string]uint32
|
||||
// flag integers so conditional fields are read iff present. The constructor's
|
||||
// unit and depth have already been charged by the caller.
|
||||
func (s *walkState) skipCtorBody(m *schemaModel, b *bin.Buffer, cl *ctorLayout, depth int) error {
|
||||
// Layer 227 constructors currently use at most flags + flags2. Keep generous fixed stack
|
||||
// storage so the allocation-free preflight remains allocation-free on the hottest flagged
|
||||
// methods; the explicit bound also prevents a future malformed/generated layout from turning
|
||||
// every request into an attacker-amplified map allocation.
|
||||
var flags [maxConstructorFlagWords]constructorFlagWord
|
||||
flagCount := 0
|
||||
for i := range cl.fields {
|
||||
f := &cl.fields[i]
|
||||
if f.isFlags {
|
||||
|
|
@ -34,59 +207,80 @@ func (m *schemaModel) skipCtorBody(b *bin.Buffer, cl *ctorLayout) error {
|
|||
if err != nil {
|
||||
return fmt.Errorf("%s.%s: %w", cl.name, f.name, err)
|
||||
}
|
||||
if flags == nil {
|
||||
flags = make(map[string]uint32, 2)
|
||||
if flagCount >= len(flags) {
|
||||
return limitf("constructor %s has more than %d flags words", cl.name, len(flags))
|
||||
}
|
||||
flags[f.name] = v
|
||||
flags[flagCount] = constructorFlagWord{name: f.name, value: v}
|
||||
flagCount++
|
||||
continue
|
||||
}
|
||||
if f.conditional() && flags[f.flagName]&(1<<uint(f.flagBit)) == 0 {
|
||||
continue
|
||||
if f.conditional() {
|
||||
var (
|
||||
flagValue uint32
|
||||
found bool
|
||||
)
|
||||
for j := 0; j < flagCount; j++ {
|
||||
if flags[j].name == f.flagName {
|
||||
flagValue = flags[j].value
|
||||
found = true
|
||||
break
|
||||
}
|
||||
}
|
||||
if !found {
|
||||
return malformedf("constructor %s conditional field %s references missing flags word %s", cl.name, f.name, f.flagName)
|
||||
}
|
||||
if flagValue&(1<<uint(f.flagBit)) == 0 {
|
||||
continue
|
||||
}
|
||||
}
|
||||
if err := m.skipValue(b, f); err != nil {
|
||||
if err := s.skipValue(m, b, f, cl, depth); err != nil {
|
||||
return fmt.Errorf("%s.%s: %w", cl.name, f.name, err)
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// skipValue advances b past one (already known-present) field value.
|
||||
func (m *schemaModel) skipValue(b *bin.Buffer, f *fieldLayout) error {
|
||||
// skipValue advances b past one already-known-present field value.
|
||||
func (s *walkState) skipValue(m *schemaModel, b *bin.Buffer, f *fieldLayout, owner *ctorLayout, depth int) error {
|
||||
switch f.kind {
|
||||
case kindInt:
|
||||
_, err := b.Int()
|
||||
return err
|
||||
case kindLong:
|
||||
_, err := b.Long()
|
||||
return err
|
||||
case kindDouble:
|
||||
_, err := b.Double()
|
||||
return err
|
||||
return skipFixed(b, 4)
|
||||
case kindLong, kindDouble:
|
||||
return skipFixed(b, 8)
|
||||
case kindInt128:
|
||||
_, err := b.Int128()
|
||||
return err
|
||||
return skipFixed(b, 16)
|
||||
case kindInt256:
|
||||
_, err := b.Int256()
|
||||
return err
|
||||
return skipFixed(b, 32)
|
||||
case kindBytes:
|
||||
_, err := b.Bytes()
|
||||
return err
|
||||
return s.skipTLBytes(b, "bytes")
|
||||
case kindString:
|
||||
_, err := b.String()
|
||||
return err
|
||||
return s.skipTLBytes(b, "string")
|
||||
case kindBool:
|
||||
_, err := b.Bool()
|
||||
return err
|
||||
if err := s.addUnits(1, "Bool constructor"); err != nil {
|
||||
return err
|
||||
}
|
||||
id, err := b.Uint32()
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if id != bin.TypeTrue && id != bin.TypeFalse {
|
||||
return malformedf("invalid Bool constructor %#08x", id)
|
||||
}
|
||||
return nil
|
||||
case kindTrue:
|
||||
return nil
|
||||
case kindVector, kindVectorBare:
|
||||
vectorDepth := depth + 1
|
||||
if vectorDepth <= 0 || vectorDepth > s.limits.maxDepth {
|
||||
return limitf("vector nesting depth %d exceeds limit %d", vectorDepth, s.limits.maxDepth)
|
||||
}
|
||||
if f.kind == kindVector {
|
||||
id, err := b.Uint32()
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if id != vectorTypeID {
|
||||
return fmt.Errorf("expected vector id, got %#08x", id)
|
||||
return malformedf("expected vector id, got %#08x", id)
|
||||
}
|
||||
}
|
||||
n, err := b.Int()
|
||||
|
|
@ -94,23 +288,146 @@ func (m *schemaModel) skipValue(b *bin.Buffer, f *fieldLayout) error {
|
|||
return err
|
||||
}
|
||||
if n < 0 {
|
||||
return fmt.Errorf("negative vector length %d", n)
|
||||
return malformedf("negative vector length %d", n)
|
||||
}
|
||||
if max := s.vectorLimit(owner, f); n > max {
|
||||
return limitf("vector %s.%s length %d exceeds limit %d", ownerName(owner), fieldName(f), n, max)
|
||||
}
|
||||
if err := s.addUnits(n, "vector "+ownerName(owner)+"."+fieldName(f)); err != nil {
|
||||
return err
|
||||
}
|
||||
if width, ok := fixedWireWidth(f.elem); ok {
|
||||
total, ok := checkedMulInt(n, width)
|
||||
if !ok {
|
||||
return malformedf("vector byte length overflow: %d * %d", n, width)
|
||||
}
|
||||
return skipFixed(b, total)
|
||||
}
|
||||
for i := 0; i < n; i++ {
|
||||
if err := m.skipValue(b, f.elem); err != nil {
|
||||
return err
|
||||
if err := s.skipValue(m, b, f.elem, nil, vectorDepth); err != nil {
|
||||
return fmt.Errorf("vector element %d: %w", i, err)
|
||||
}
|
||||
}
|
||||
return nil
|
||||
case kindObject:
|
||||
return m.skipObject(b)
|
||||
return s.skipObject(m, b, depth+1)
|
||||
case kindBareObject:
|
||||
bareDepth := depth + 1
|
||||
if err := s.enter(bareDepth, "bare constructor"); err != nil {
|
||||
return err
|
||||
}
|
||||
cl, ok := m.bareByT[f.typeName]
|
||||
if !ok {
|
||||
return fmt.Errorf("unknown bare type %q", f.typeName)
|
||||
return malformedf("unknown bare type %q", f.typeName)
|
||||
}
|
||||
return m.skipCtorBody(b, cl)
|
||||
return s.skipCtorBody(m, b, cl, bareDepth)
|
||||
default:
|
||||
return fmt.Errorf("bad wire kind %d", f.kind)
|
||||
return malformedf("bad wire kind %d", f.kind)
|
||||
}
|
||||
}
|
||||
|
||||
// skipTLBytes parses TL's 1/4-byte length prefix directly and advances the
|
||||
// input slice. Unlike bin.Buffer.Bytes it never copies payload data.
|
||||
func (s *walkState) skipTLBytes(b *bin.Buffer, what string) error {
|
||||
if len(b.Buf) == 0 {
|
||||
return io.ErrUnexpectedEOF
|
||||
}
|
||||
var header, payload uint64
|
||||
switch b.Buf[0] {
|
||||
case 254:
|
||||
if len(b.Buf) < 4 {
|
||||
return io.ErrUnexpectedEOF
|
||||
}
|
||||
header = 4
|
||||
payload = uint64(b.Buf[1]) | uint64(b.Buf[2])<<8 | uint64(b.Buf[3])<<16
|
||||
case 255:
|
||||
return malformedf("invalid %s length prefix 255", what)
|
||||
default:
|
||||
header = 1
|
||||
payload = uint64(b.Buf[0])
|
||||
}
|
||||
if err := s.addBytes(payload, what); err != nil {
|
||||
return err
|
||||
}
|
||||
encoded, ok := checkedAddUint64(header, payload)
|
||||
if !ok {
|
||||
return malformedf("%s encoded length overflow", what)
|
||||
}
|
||||
withPadding, ok := checkedAddUint64(encoded, 3)
|
||||
if !ok {
|
||||
return malformedf("%s padded length overflow", what)
|
||||
}
|
||||
padded := withPadding &^ uint64(3)
|
||||
if padded > uint64(math.MaxInt) {
|
||||
return malformedf("%s padded length %d overflows int", what, padded)
|
||||
}
|
||||
if uint64(len(b.Buf)) < padded {
|
||||
return io.ErrUnexpectedEOF
|
||||
}
|
||||
b.Buf = b.Buf[int(padded):]
|
||||
return nil
|
||||
}
|
||||
|
||||
func skipFixed(b *bin.Buffer, n int) error {
|
||||
if n < 0 {
|
||||
return malformedf("negative fixed-width skip %d", n)
|
||||
}
|
||||
if len(b.Buf) < n {
|
||||
return io.ErrUnexpectedEOF
|
||||
}
|
||||
b.Buf = b.Buf[n:]
|
||||
return nil
|
||||
}
|
||||
|
||||
func fixedWireWidth(f *fieldLayout) (int, bool) {
|
||||
if f == nil {
|
||||
return 0, false
|
||||
}
|
||||
switch f.kind {
|
||||
case kindInt:
|
||||
return 4, true
|
||||
case kindLong, kindDouble:
|
||||
return 8, true
|
||||
case kindInt128:
|
||||
return 16, true
|
||||
case kindInt256:
|
||||
return 32, true
|
||||
case kindTrue:
|
||||
return 0, true
|
||||
default:
|
||||
// Bool deliberately stays on the element loop so constructor ids are
|
||||
// validated and charged to the aggregate constructor budget.
|
||||
return 0, false
|
||||
}
|
||||
}
|
||||
|
||||
func checkedMulInt(a, b int) (int, bool) {
|
||||
if a < 0 || b < 0 {
|
||||
return 0, false
|
||||
}
|
||||
if a != 0 && b > math.MaxInt/a {
|
||||
return 0, false
|
||||
}
|
||||
return a * b, true
|
||||
}
|
||||
|
||||
func checkedAddUint64(a, b uint64) (uint64, bool) {
|
||||
if b > math.MaxUint64-a {
|
||||
return 0, false
|
||||
}
|
||||
return a + b, true
|
||||
}
|
||||
|
||||
func ownerName(cl *ctorLayout) string {
|
||||
if cl == nil || cl.name == "" {
|
||||
return "<nested>"
|
||||
}
|
||||
return cl.name
|
||||
}
|
||||
|
||||
func fieldName(f *fieldLayout) string {
|
||||
if f == nil || f.name == "" {
|
||||
return "<element>"
|
||||
}
|
||||
return f.name
|
||||
}
|
||||
|
|
|
|||
263
internal/compat/layerwire/walk_limits_test.go
Normal file
263
internal/compat/layerwire/walk_limits_test.go
Normal file
|
|
@ -0,0 +1,263 @@
|
|||
package layerwire
|
||||
|
||||
import (
|
||||
"errors"
|
||||
"math"
|
||||
"testing"
|
||||
|
||||
"github.com/gotd/td/bin"
|
||||
"github.com/gotd/td/tg"
|
||||
)
|
||||
|
||||
func TestValidateCanonicalRequestFlaggedHotPathAllocatesNothing(t *testing.T) {
|
||||
var body bin.Buffer
|
||||
req := &tg.MessagesSendMessageRequest{
|
||||
Peer: &tg.InputPeerSelf{},
|
||||
Message: "hello",
|
||||
RandomID: 7,
|
||||
}
|
||||
if err := req.Encode(&body); err != nil {
|
||||
t.Fatalf("encode request: %v", err)
|
||||
}
|
||||
if err := ValidateCanonicalRequest(body.Buf); err != nil {
|
||||
t.Fatalf("validate request: %v", err)
|
||||
}
|
||||
if allocs := testing.AllocsPerRun(1000, func() {
|
||||
if err := ValidateCanonicalRequest(body.Buf); err != nil {
|
||||
panic(err)
|
||||
}
|
||||
}); allocs != 0 {
|
||||
t.Fatalf("canonical request preflight allocations = %.2f, want 0", allocs)
|
||||
}
|
||||
}
|
||||
|
||||
func TestValidateCanonicalRequestVectorLimits(t *testing.T) {
|
||||
editCloseFriends := canonical.byName["contacts.editCloseFriends"]
|
||||
if editCloseFriends == nil {
|
||||
t.Fatal("contacts.editCloseFriends missing from canonical schema")
|
||||
}
|
||||
|
||||
t.Run("explicit_5000_override", func(t *testing.T) {
|
||||
var body bin.Buffer
|
||||
body.PutID(editCloseFriends.crc)
|
||||
body.PutVectorHeader(5000)
|
||||
for i := 0; i < 5000; i++ {
|
||||
body.PutLong(int64(i))
|
||||
}
|
||||
if err := ValidateCanonicalRequest(body.Buf); err != nil {
|
||||
t.Fatalf("validate legal 5000-element close-friends request: %v", err)
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("override_stops_at_5000", func(t *testing.T) {
|
||||
var body bin.Buffer
|
||||
body.PutID(editCloseFriends.crc)
|
||||
body.PutVectorHeader(5001)
|
||||
err := ValidateCanonicalRequest(body.Buf)
|
||||
if !errors.Is(err, ErrResourceLimit) {
|
||||
t.Fatalf("error = %v, want ErrResourceLimit", err)
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("default_4096", func(t *testing.T) {
|
||||
getMessages := canonical.byName["messages.getMessages"]
|
||||
var body bin.Buffer
|
||||
body.PutID(getMessages.crc)
|
||||
body.PutVectorHeader(defaultMaxVectorElements + 1)
|
||||
err := ValidateCanonicalRequest(body.Buf)
|
||||
if !errors.Is(err, ErrResourceLimit) {
|
||||
t.Fatalf("error = %v, want ErrResourceLimit", err)
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("max_int32_count_rejected_before_iteration", func(t *testing.T) {
|
||||
var body bin.Buffer
|
||||
body.PutID(editCloseFriends.crc)
|
||||
body.PutID(vectorTypeID)
|
||||
body.PutInt32(math.MaxInt32)
|
||||
err := ValidateCanonicalRequest(body.Buf)
|
||||
if !errors.Is(err, ErrResourceLimit) {
|
||||
t.Fatalf("error = %v, want ErrResourceLimit", err)
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
func TestValidateCanonicalRequestDepthLimit(t *testing.T) {
|
||||
invoke := canonical.byName["invokeWithoutUpdates"]
|
||||
leaf := canonical.byName["help.getConfig"]
|
||||
if invoke == nil || leaf == nil {
|
||||
t.Fatal("generic wrapper methods missing from canonical schema")
|
||||
}
|
||||
request := func(wrappers int) []byte {
|
||||
var body bin.Buffer
|
||||
for i := 0; i < wrappers; i++ {
|
||||
body.PutID(invoke.crc)
|
||||
}
|
||||
body.PutID(leaf.crc)
|
||||
return body.Buf
|
||||
}
|
||||
if err := ValidateCanonicalRequest(request(defaultMaxWalkDepth - 1)); err != nil {
|
||||
t.Fatalf("depth exactly %d rejected: %v", defaultMaxWalkDepth, err)
|
||||
}
|
||||
err := ValidateCanonicalRequest(request(defaultMaxWalkDepth))
|
||||
if !errors.Is(err, ErrResourceLimit) {
|
||||
t.Fatalf("depth %d error = %v, want ErrResourceLimit", defaultMaxWalkDepth+1, err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestTLBytesSkipIsZeroCopyAndBounded(t *testing.T) {
|
||||
var encoded bin.Buffer
|
||||
encoded.PutBytes([]byte("payload"))
|
||||
fieldLen := len(encoded.Buf)
|
||||
raw := append(encoded.Copy(), 0xaa, 0xbb, 0xcc, 0xdd)
|
||||
b := &bin.Buffer{Buf: raw}
|
||||
walk := newWalkState()
|
||||
if err := walk.skipTLBytes(b, "bytes"); err != nil {
|
||||
t.Fatalf("skip bytes: %v", err)
|
||||
}
|
||||
if len(b.Buf) != 4 || &b.Buf[0] != &raw[fieldLen] {
|
||||
t.Fatalf("walker did not retain the original backing buffer")
|
||||
}
|
||||
|
||||
t.Run("per_field_budget", func(t *testing.T) {
|
||||
limited := newWalkState()
|
||||
limited.limits.maxFieldBytes = 3
|
||||
probe := &bin.Buffer{Buf: encoded.Copy()}
|
||||
err := limited.skipTLBytes(probe, "bytes")
|
||||
if !errors.Is(err, ErrResourceLimit) {
|
||||
t.Fatalf("error = %v, want ErrResourceLimit", err)
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("aggregate_budget", func(t *testing.T) {
|
||||
limited := newWalkState()
|
||||
limited.limits.maxTotalBytes = 10
|
||||
first := &bin.Buffer{Buf: encoded.Copy()}
|
||||
if err := limited.skipTLBytes(first, "bytes"); err != nil {
|
||||
t.Fatalf("first field: %v", err)
|
||||
}
|
||||
second := &bin.Buffer{Buf: encoded.Copy()}
|
||||
err := limited.skipTLBytes(second, "bytes")
|
||||
if !errors.Is(err, ErrResourceLimit) {
|
||||
t.Fatalf("second error = %v, want ErrResourceLimit", err)
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("truncated_payload_is_malformed", func(t *testing.T) {
|
||||
importAuth := canonical.byName["auth.importAuthorization"]
|
||||
var body bin.Buffer
|
||||
body.PutID(importAuth.crc)
|
||||
body.PutLong(1)
|
||||
body.Put([]byte{5, 'a', 'b'}) // declares five bytes, lacks payload/padding
|
||||
err := ValidateCanonicalRequest(body.Buf)
|
||||
if !errors.Is(err, ErrMalformed) || errors.Is(err, ErrResourceLimit) {
|
||||
t.Fatalf("error = %v, want only ErrMalformed", err)
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
func TestInboundTransformsShareWalkerBudgets(t *testing.T) {
|
||||
t.Run("canonical_alias", func(t *testing.T) {
|
||||
var body bin.Buffer
|
||||
body.PutID(0x41d41ade) // DrKLO messages.forwardMessages alias
|
||||
body.PutUint32(0)
|
||||
body.PutID(canonical.byName["inputPeerEmpty"].crc)
|
||||
body.PutID(vectorTypeID)
|
||||
body.PutInt32(math.MaxInt32)
|
||||
_, ok, err := UpgradeInbound(0x41d41ade, &body)
|
||||
if !ok || !errors.Is(err, ErrResourceLimit) {
|
||||
t.Fatalf("ok=%v error=%v, want matched ErrResourceLimit", ok, err)
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("drift_body_transform", func(t *testing.T) {
|
||||
var body bin.Buffer
|
||||
body.PutID(0x2e1ee318) // DrKLO langpack.getStrings body transform
|
||||
body.PutString("en")
|
||||
body.PutID(vectorTypeID)
|
||||
body.PutInt32(math.MaxInt32)
|
||||
_, ok, err := UpgradeInbound(0x2e1ee318, &body)
|
||||
if !ok || !errors.Is(err, ErrResourceLimit) {
|
||||
t.Fatalf("ok=%v error=%v, want matched ErrResourceLimit", ok, err)
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("outbound_structural_transform", func(t *testing.T) {
|
||||
poll := canonical.byName["pollAnswerVoters"]
|
||||
var body bin.Buffer
|
||||
body.PutID(poll.crc)
|
||||
body.PutUint32(1 << 2)
|
||||
body.PutBytes(nil)
|
||||
body.PutInt(1)
|
||||
body.PutID(vectorTypeID)
|
||||
body.PutInt32(math.MaxInt32)
|
||||
_, err := Transcode(body.Buf, CanonicalLayer-1)
|
||||
if !errors.Is(err, ErrResourceLimit) {
|
||||
t.Fatalf("error = %v, want ErrResourceLimit", err)
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
func TestWalkerArithmeticAndMalformedClassification(t *testing.T) {
|
||||
if defaultMaxWalkUnits != 131072 {
|
||||
t.Fatalf("default constructor/vector budget = %d, want 131072", defaultMaxWalkUnits)
|
||||
}
|
||||
if _, ok := checkedMulInt(math.MaxInt, 2); ok {
|
||||
t.Fatal("checkedMulInt accepted overflow")
|
||||
}
|
||||
if _, ok := checkedAddUint64(math.MaxUint64, 1); ok {
|
||||
t.Fatal("checkedAddUint64 accepted overflow")
|
||||
}
|
||||
|
||||
t.Run("aggregate_constructor_and_vector_units", func(t *testing.T) {
|
||||
editCloseFriends := canonical.byName["contacts.editCloseFriends"]
|
||||
var body bin.Buffer
|
||||
body.PutID(editCloseFriends.crc)
|
||||
body.PutVectorHeader(4)
|
||||
for i := 0; i < 4; i++ {
|
||||
body.PutLong(int64(i))
|
||||
}
|
||||
walk := newWalkState()
|
||||
walk.limits.maxUnits = 4 // top constructor + four elements needs five
|
||||
probe := &bin.Buffer{Buf: body.Buf}
|
||||
err := walk.skipObject(canonical, probe, 1)
|
||||
if !errors.Is(err, ErrResourceLimit) {
|
||||
t.Fatalf("error = %v, want ErrResourceLimit", err)
|
||||
}
|
||||
})
|
||||
|
||||
tests := []struct {
|
||||
name string
|
||||
body []byte
|
||||
}{
|
||||
{name: "empty"},
|
||||
{name: "unknown_constructor", body: []byte{1, 2, 3, 4}},
|
||||
{name: "trailing_bytes", body: append(methodIDBytes(canonical.byName["help.getConfig"].crc), 0, 0, 0, 0)},
|
||||
}
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
err := ValidateCanonicalRequest(tt.body)
|
||||
if !errors.Is(err, ErrMalformed) || errors.Is(err, ErrResourceLimit) {
|
||||
t.Fatalf("error = %v, want only ErrMalformed", err)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func methodIDBytes(id uint32) []byte {
|
||||
var b bin.Buffer
|
||||
b.PutID(id)
|
||||
return b.Buf
|
||||
}
|
||||
|
||||
func FuzzValidateCanonicalRequest(f *testing.F) {
|
||||
f.Add(methodIDBytes(canonical.byName["help.getConfig"].crc))
|
||||
f.Add([]byte{})
|
||||
f.Add([]byte{1, 2, 3, 4})
|
||||
f.Fuzz(func(t *testing.T, body []byte) {
|
||||
err := ValidateCanonicalRequest(body)
|
||||
if err != nil && !errors.Is(err, ErrMalformed) && !errors.Is(err, ErrResourceLimit) {
|
||||
t.Fatalf("unclassified walker error: %v", err)
|
||||
}
|
||||
})
|
||||
}
|
||||
Loading…
Add table
Add a link
Reference in a new issue