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
20 changes: 20 additions & 0 deletions context.go
Original file line number Diff line number Diff line change
Expand Up @@ -5,6 +5,7 @@ package echo

import (
"bytes"
stdContext "context"
"encoding/xml"
"errors"
"fmt"
Expand Down Expand Up @@ -460,6 +461,10 @@ func (c *Context) Bind(i any) error {
return c.echo.Binder.Bind(c, i)
}

type validatorCtx interface {
ValidateCtx(ctx stdContext.Context, i any) error
}

// Validate validates provided `i`. It is usually called after `Context#Bind()`.
// Validator must be registered using `Echo#Validator`.
func (c *Context) Validate(i any) error {
Expand All @@ -469,6 +474,21 @@ func (c *Context) Validate(i any) error {
return c.echo.Validator.Validate(i)
}

// ValidateCtx validates i using the current request's context when the registered
// Validator implements ValidateCtx(context.Context, any) error. Otherwise it calls Validate.
// The validator must still implement Validate to satisfy the Validator interface.
// If there is no request, context.Background() is used.
func (c *Context) ValidateCtx(i any) error {
if v, ok := c.echo.Validator.(validatorCtx); ok {
ctx := stdContext.Background()
if req := c.Request(); req != nil {
ctx = req.Context()
}
return v.ValidateCtx(ctx, i)
}
return c.Validate(i)
}

// Render renders a template with data and sends a text/html response with status
// code. Renderer must be registered using `Echo.Renderer`.
func (c *Context) Render(code int, name string, data any) (err error) {
Expand Down
103 changes: 102 additions & 1 deletion context_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -5,9 +5,11 @@ package echo

import (
"bytes"
stdContext "context"
"crypto/tls"
"encoding/json"
"encoding/xml"
"errors"
"fmt"
"io"
"io/fs"
Expand Down Expand Up @@ -1119,10 +1121,109 @@ func TestContext_Validate(t *testing.T) {
e := New()
c := e.NewContext(nil, nil)

assert.Error(t, c.Validate(struct{}{}))
assert.ErrorIs(t, c.Validate(struct{}{}), ErrValidatorNotRegistered)
assert.ErrorIs(t, c.ValidateCtx(struct{}{}), ErrValidatorNotRegistered)

e.Validator = &validator{}
assert.NoError(t, c.Validate(struct{}{}))
assert.NoError(t, c.ValidateCtx(struct{}{}))
}

type validationFunc func(any) error

func (v validationFunc) Validate(i any) error { return v(i) }

type contextValidator struct {
validationFunc
validateCtx func(stdContext.Context, any) error
}

func (v contextValidator) ValidateCtx(ctx stdContext.Context, i any) error {
return v.validateCtx(ctx, i)
}

func TestContext_Validate_legacyError(t *testing.T) {
e := New()
c := e.NewContext(nil, nil)
payload := &struct{ Name string }{Name: "Jon Snow"}
wantErr := errors.New("validation failed")
e.Validator = validationFunc(func(i any) error {
assert.Same(t, payload, i)
return wantErr
})
assert.Same(t, wantErr, c.Validate(payload))
assert.Same(t, wantErr, c.ValidateCtx(payload))
}

type validationWrapper struct {
contextValidator
err error
}

func (v validationWrapper) Validate(any) error { return v.err }

func TestContext_Validate_wrapper(t *testing.T) {
e := New()
c := e.NewContext(nil, nil)
wantErr := errors.New("wrapper validation failed")
e.Validator = validationWrapper{
contextValidator: contextValidator{
validateCtx: func(stdContext.Context, any) error {
t.Fatal("Validate must not call an embedded ValidateCtx method")
return nil
},
},
err: wantErr,
}
assert.Same(t, wantErr, c.Validate(struct{}{}))
}

func TestContext_ValidateCtx_withContext(t *testing.T) {
ctx, cancel := stdContext.WithCancel(stdContext.Background())
defer cancel()
e := New()
c := e.NewContext(nil, nil)
c.SetRequest(httptest.NewRequest(http.MethodGet, "/", nil).WithContext(ctx))
payload := &struct{ Name string }{Name: "Jon Snow"}
wantErr := errors.New("validation failed")
calls := 0
e.Validator = contextValidator{
validationFunc: func(any) error {
t.Fatal("Validate must not be called when ValidateCtx is implemented")
return nil
},
validateCtx: func(got stdContext.Context, i any) error {
calls++
assert.Same(t, ctx, got)
assert.Same(t, payload, i)
return wantErr
},
}

assert.Same(t, wantErr, c.ValidateCtx(payload))
assert.Equal(t, 1, calls)

cancel()
assert.Same(t, wantErr, c.ValidateCtx(payload))
assert.Equal(t, 2, calls)
}

func TestContext_ValidateCtx_withoutRequest(t *testing.T) {
e := New()
c := e.NewContext(nil, nil)
var gotCtx stdContext.Context
e.Validator = contextValidator{
validateCtx: func(ctx stdContext.Context, i any) error {
gotCtx = ctx
return nil
},
}

assert.NoError(t, c.ValidateCtx(struct{}{}))
if assert.NotNil(t, gotCtx) {
assert.NoError(t, gotCtx.Err())
assert.Nil(t, gotCtx.Done())
}
}

func TestContext_QueryString(t *testing.T) {
Expand Down
Loading