diff --git a/render.go b/render.go index 75a90e2..44db049 100644 --- a/render.go +++ b/render.go @@ -90,40 +90,99 @@ func renderer(w http.ResponseWriter, r *http.Request, v Renderer) error { return nil } +func asBinder(f reflect.Value) (Binder, bool) { + if !f.IsValid() { + return nil, false + } + switch f.Kind() { + case reflect.Ptr, reflect.Interface: + if f.IsNil() { + return nil, false + } + } + if f.Type().Implements(binderType) { + if !f.CanInterface() { + return nil, false + } + b, ok := f.Interface().(Binder) + return b, ok + } + if f.CanAddr() && f.Addr().Type().Implements(binderType) { + b, ok := f.Addr().Interface().(Binder) + return b, ok + } + if f.CanInterface() && f.Kind() != reflect.Ptr && f.Kind() != reflect.Interface { + ptr := reflect.New(f.Type()) + ptr.Elem().Set(f) + if ptr.Type().Implements(binderType) { + b, ok := ptr.Interface().(Binder) + return b, ok + } + } + return nil, false +} + +func bindValue(r *http.Request, f reflect.Value) error { + if !f.IsValid() { + return nil + } + if b, ok := asBinder(f); ok { + return binder(r, b) + } + switch f.Kind() { + case reflect.Ptr, reflect.Interface: + if f.IsNil() { + return nil + } + return bindValue(r, f.Elem()) + case reflect.Slice, reflect.Array: + for i := 0; i < f.Len(); i++ { + if err := bindValue(r, f.Index(i)); err != nil { + return err + } + } + case reflect.Map: + for _, key := range f.MapKeys() { + if err := bindValue(r, f.MapIndex(key)); err != nil { + return err + } + } + } + return nil +} + // Executed bottom-up func binder(r *http.Request, v Binder) error { rv := reflect.ValueOf(v) if rv.Kind() == reflect.Ptr { + if rv.IsNil() { + return v.Bind(r) + } rv = rv.Elem() } - // Call Binder on non-struct types right away - if rv.Kind() != reflect.Struct { - return v.Bind(r) - } - - // For structs, we call Bind on each field that implements Binder - for i := 0; i < rv.NumField(); i++ { - f := rv.Field(i) - if f.Type().Implements(binderType) { - - if isNil(f) { - continue + switch rv.Kind() { + case reflect.Struct: + for i := 0; i < rv.NumField(); i++ { + if err := bindValue(r, rv.Field(i)); err != nil { + return err } - - fv := f.Interface().(Binder) - if err := binder(r, fv); err != nil { + } + case reflect.Slice, reflect.Array: + for i := 0; i < rv.Len(); i++ { + if err := bindValue(r, rv.Index(i)); err != nil { + return err + } + } + case reflect.Map: + for _, key := range rv.MapKeys() { + if err := bindValue(r, rv.MapIndex(key)); err != nil { return err } } } - // We call it bottom-up - if err := v.Bind(r); err != nil { - return err - } - - return nil + return v.Bind(r) } var ( diff --git a/render_bind_test.go b/render_bind_test.go new file mode 100644 index 0000000..3a37dbf --- /dev/null +++ b/render_bind_test.go @@ -0,0 +1,135 @@ +package render + +import ( + "errors" + "net/http" + "net/http/httptest" + "strings" + "testing" +) + +type bindBody struct { + Name string `json:"name"` + Children Children `json:"children"` +} + +func (b *bindBody) Bind(r *http.Request) error { + if b.Name == "" { + return errors.New("empty body name") + } + return nil +} + +type Children []Child + +func (c Children) Bind(r *http.Request) error { + if c == nil { + return errors.New("no children provided") + } + return nil +} + +type Child struct { + Name string `json:"name"` +} + +func (c *Child) Bind(r *http.Request) error { + if c.Name == "" { + return errors.New("empty child name") + } + return nil +} + +type bindPlainSlice struct { + Kids []Child `json:"kids"` +} + +func (b *bindPlainSlice) Bind(r *http.Request) error { return nil } + +type translatedField map[string]string + +func (t translatedField) Bind(r *http.Request) error { + if len(t) == 0 { + return errors.New("empty translation") + } + return nil +} + +type bindMapField struct { + Title translatedField `json:"title"` +} + +func (b *bindMapField) Bind(r *http.Request) error { return nil } + +type mapValueChild struct { + Items map[string]Child `json:"items"` +} + +func (m *mapValueChild) Bind(r *http.Request) error { return nil } + +func TestBindWalksSliceElements(t *testing.T) { + r := httptest.NewRequest(http.MethodPost, "/", strings.NewReader( + `{"name":"ok","children":[{"name":"a"},{"name":""}]}`, + )) + r.Header.Set("Content-Type", "application/json") + var body bindBody + err := Bind(r, &body) + if err == nil || err.Error() != "empty child name" { + t.Fatalf("got %v, want empty child name", err) + } +} + +func TestBindNamedSliceType(t *testing.T) { + r := httptest.NewRequest(http.MethodPost, "/", strings.NewReader( + `{"name":"ok","children":[{"name":"a"}]}`, + )) + r.Header.Set("Content-Type", "application/json") + var body bindBody + if err := Bind(r, &body); err != nil { + t.Fatal(err) + } +} + +func TestBindNilNamedSlice(t *testing.T) { + r := httptest.NewRequest(http.MethodPost, "/", strings.NewReader(`{"name":"ok"}`)) + r.Header.Set("Content-Type", "application/json") + var body bindBody + err := Bind(r, &body) + if err == nil || err.Error() != "no children provided" { + t.Fatalf("got %v, want no children provided", err) + } +} + +func TestBindUnnamedSliceElements(t *testing.T) { + r := httptest.NewRequest(http.MethodPost, "/", strings.NewReader( + `{"kids":[{"name":""}]}`, + )) + r.Header.Set("Content-Type", "application/json") + var body bindPlainSlice + err := Bind(r, &body) + if err == nil || err.Error() != "empty child name" { + t.Fatalf("got %v, want empty child name", err) + } +} + +func TestBindMapTypeField(t *testing.T) { + r := httptest.NewRequest(http.MethodPost, "/", strings.NewReader(`{"title":{}}`)) + r.Header.Set("Content-Type", "application/json") + var body bindMapField + err := Bind(r, &body) + if err == nil || err.Error() != "empty translation" { + t.Fatalf("got %v, want empty translation", err) + } +} + +func TestBindMapValues(t *testing.T) { + r := httptest.NewRequest(http.MethodPost, "/", strings.NewReader( + `{"items":{"a":{"name":""}}}`, + )) + r.Header.Set("Content-Type", "application/json") + var body mapValueChild + err := Bind(r, &body) + if err == nil || err.Error() != "empty child name" { + t.Fatalf("got %v, want empty child name", err) + } +}