@@ -3,13 +3,16 @@ package croupier
33
44import (
55 "context"
6+ "encoding/json"
67 "fmt"
8+ "strings"
79 "sync"
810 "sync/atomic"
911 "time"
1012
1113 agentv1 "github.com/cuihairu/croupier/sdks/go/pkg/pb/croupier/agent/v1"
1214 sdkv1 "github.com/cuihairu/croupier/sdks/go/pkg/pb/croupier/sdk/v1"
15+ "github.com/santhosh-tekuri/jsonschema/v6"
1316 "google.golang.org/protobuf/proto"
1417
1518 "github.com/cuihairu/croupier/sdks/go/pkg/croupier/protocol"
@@ -496,6 +499,16 @@ func (h *tcpRPCHandler) invoke(ctx context.Context, msgID uint32, reqID uint32,
496499 return nil , fmt .Errorf ("function not found: %s" , req .FunctionId )
497500 }
498501
502+ // 入站 payload 校验(按函数声明的 input schema):失败回错误 payload,
503+ // 游戏逻辑不会看到非法输入(服务端仍是权威校验方)。
504+ if err := h .manager .validateInboundPayload (req .FunctionId , req .Payload ); err != nil {
505+ errResp := & sdkv1.InvokeResponse {Payload : []byte (`{"error":` + jsonString (err .Error ()) + `}` )}
506+ if b , marshalErr := proto .Marshal (errResp ); marshalErr == nil {
507+ return b , nil
508+ }
509+ return nil , err
510+ }
511+
499512 // OTel 一期传播:metadata trace 字段进 handler 上下文(零侵入,无则原 ctx)
500513 ctx = WithTraceMetadata (ctx , req .GetMetadata ())
501514 result , err := handler (ctx , req .Payload )
@@ -509,6 +522,62 @@ func (h *tcpRPCHandler) invoke(ctx context.Context, msgID uint32, reqID uint32,
509522 return proto .Marshal (resp )
510523}
511524
525+ // jsonString 序列化为 JSON 字符串字面量(含引号),用于拼接错误 payload。
526+ func jsonString (s string ) string {
527+ b , err := json .Marshal (s )
528+ if err != nil {
529+ return `""`
530+ }
531+ return string (b )
532+ }
533+
534+ // validateInboundPayload 按 ProviderFunctionDescriptor 声明的 input schema
535+ // 校验入站 payload。未开启开关 / 未找到描述符 / schema 为空 / schema 或
536+ // payload 非法 JSON 时跳过(服务端仍是权威校验方);编译结果按函数缓存。
537+ func (m * TCPManager ) validateInboundPayload (functionID string , payload []byte ) error {
538+ if ! m .config .ValidateInputPayloads {
539+ return nil
540+ }
541+ m .mu .RLock ()
542+ var schemaRaw string
543+ for _ , descriptor := range m .functions {
544+ if descriptor != nil && descriptor .Id == functionID {
545+ schemaRaw = descriptor .InputSchema
546+ break
547+ }
548+ }
549+ m .mu .RUnlock ()
550+ if strings .TrimSpace (schemaRaw ) == "" {
551+ return nil
552+ }
553+
554+ var schemaDoc interface {}
555+ if err := json .Unmarshal ([]byte (schemaRaw ), & schemaDoc ); err != nil {
556+ // 非法 schema 不在 provider 侧报错(注册校验负责),跳过
557+ return nil
558+ }
559+
560+ var value interface {}
561+ if err := json .Unmarshal (payload , & value ); err != nil {
562+ return fmt .Errorf ("payload must be valid JSON: %w" , err )
563+ }
564+
565+ compiler := jsonschema .NewCompiler ()
566+ compiler .DefaultDraft (jsonschema .Draft7 )
567+ if err := compiler .AddResource ("schema.json" , schemaDoc ); err != nil {
568+ // 编译失败视为 schema 缺陷,跳过(与非法 schema 同策略)
569+ return nil
570+ }
571+ sch , err := compiler .Compile ("schema.json" )
572+ if err != nil {
573+ return nil
574+ }
575+ if err := sch .Validate (value ); err != nil {
576+ return fmt .Errorf ("payload validation failed: %s" , err .Error ())
577+ }
578+ return nil
579+ }
580+
512581func (h * tcpRPCHandler ) startTask (ctx context.Context , msgID uint32 , reqID uint32 , body []byte ) (respBody []byte , err error ) {
513582 req := & sdkv1.InvokeRequest {}
514583 if err := proto .Unmarshal (body , req ); err != nil {
@@ -523,6 +592,11 @@ func (h *tcpRPCHandler) startTask(ctx context.Context, msgID uint32, reqID uint3
523592 return nil , fmt .Errorf ("function not found: %s" , req .FunctionId )
524593 }
525594
595+ // 入站 payload 校验(同 invoke):失败回错误给 agent,任务不启动。
596+ if err := h .manager .validateInboundPayload (req .FunctionId , req .Payload ); err != nil {
597+ return nil , err
598+ }
599+
526600 // OTel 一期传播:任务 handler 上下文同样可读 trace 字段
527601 taskCtx , cancel := context .WithCancel (WithTraceMetadata (ctx , req .GetMetadata ()))
528602
0 commit comments