- Date (DATE) and TimeOfDay (TIME) with Scanner, Valuer, JSON and text forms - Fill falls back to encoding.TextUnmarshaler for string sources after every existing conversion, so time.Time and the new types fill from JSON - required treats a zero time.Time, Date or TimeOfDay as empty - lagoon README, models and casts-and-validation docs
281 lines
8.1 KiB
Go
281 lines
8.1 KiB
Go
package lagoon
|
|
|
|
import (
|
|
"database/sql"
|
|
"encoding"
|
|
"encoding/json"
|
|
"fmt"
|
|
"log/slog"
|
|
"reflect"
|
|
"strconv"
|
|
"strings"
|
|
"sync"
|
|
)
|
|
|
|
// HasFillable is the Go form of Eloquent $fillable: the model's backstop
|
|
// allow-list for mass assignment (D-05).
|
|
type HasFillable interface {
|
|
Fillable() []string
|
|
}
|
|
|
|
// HasHidden is the Go form of Eloquent $hidden: column names that must not
|
|
// appear in an accidental JSON marshal (D-08).
|
|
type HasHidden interface {
|
|
Hidden() []string
|
|
}
|
|
|
|
var droppedKeys sync.Map
|
|
|
|
// FillTypeError reports a requested value that cannot be stored in its
|
|
// column: a fraction, an exponent or an overflow for an integer field, or a
|
|
// value of the wrong type. Key is the column the value was requested for, so
|
|
// a caller can answer it as a validation failure on that field. Fill returns
|
|
// any other failure, such as a model that is not a struct pointer, as a
|
|
// plain error.
|
|
type FillTypeError struct {
|
|
Key string
|
|
Err error
|
|
}
|
|
|
|
func (e *FillTypeError) Error() string { return "lagoon: fill " + e.Key + ": " + e.Err.Error() }
|
|
|
|
func (e *FillTypeError) Unwrap() error { return e.Err }
|
|
|
|
// Fill copies requested keys onto model only when they are also in allowed.
|
|
// Unknown and non-fillable keys are dropped with no error (D-06). In
|
|
// non-production, each type+key pair is logged once. A value that does not
|
|
// fit its column is a *FillTypeError.
|
|
func Fill(model any, allowed []string, requested map[string]any, production bool) error {
|
|
if model == nil {
|
|
return fmt.Errorf("lagoon: fill model is nil")
|
|
}
|
|
rv := reflect.ValueOf(model)
|
|
if rv.Kind() != reflect.Ptr || rv.IsNil() {
|
|
return fmt.Errorf("lagoon: fill model must be a non-nil pointer")
|
|
}
|
|
rv = rv.Elem()
|
|
if rv.Kind() != reflect.Struct {
|
|
return fmt.Errorf("lagoon: fill model must point to a struct")
|
|
}
|
|
rt := rv.Type()
|
|
typeName := rt.String()
|
|
for key, val := range requested {
|
|
if !allowListed(key, allowed) {
|
|
logDroppedKeyOnce(production, typeName, key)
|
|
continue
|
|
}
|
|
field, ok := fieldByColumn(rv, rt, key)
|
|
if !ok {
|
|
logDroppedKeyOnce(production, typeName, key)
|
|
continue
|
|
}
|
|
if err := setField(field, val); err != nil {
|
|
return &FillTypeError{Key: key, Err: err}
|
|
}
|
|
}
|
|
return nil
|
|
}
|
|
|
|
func logDroppedKeyOnce(production bool, typeName, key string) {
|
|
if production {
|
|
return
|
|
}
|
|
k := typeName + "." + key
|
|
if _, loaded := droppedKeys.LoadOrStore(k, struct{}{}); loaded {
|
|
return
|
|
}
|
|
slog.Warn("lagoon: dropped non-fillable key", "type", typeName, "key", key)
|
|
}
|
|
|
|
func fieldByColumn(rv reflect.Value, rt reflect.Type, column string) (reflect.Value, bool) {
|
|
for i := 0; i < rt.NumField(); i++ {
|
|
f := rt.Field(i)
|
|
if !f.IsExported() {
|
|
continue
|
|
}
|
|
if gormColumn(f.Tag.Get("gorm")) == column {
|
|
return rv.Field(i), true
|
|
}
|
|
}
|
|
return reflect.Value{}, false
|
|
}
|
|
|
|
func gormColumn(tag string) string {
|
|
for _, part := range strings.Split(tag, ";") {
|
|
part = strings.TrimSpace(part)
|
|
if after, ok := strings.CutPrefix(part, "column:"); ok {
|
|
return after
|
|
}
|
|
}
|
|
return ""
|
|
}
|
|
|
|
func setField(field reflect.Value, val any) error {
|
|
if !field.CanSet() {
|
|
return fmt.Errorf("field cannot be set")
|
|
}
|
|
if val == nil {
|
|
field.Set(reflect.Zero(field.Type()))
|
|
if field.CanAddr() {
|
|
if scanner, ok := field.Addr().Interface().(sql.Scanner); ok {
|
|
return scanner.Scan(nil)
|
|
}
|
|
}
|
|
return nil
|
|
}
|
|
src := reflect.ValueOf(val)
|
|
if field.Kind() == reflect.Ptr {
|
|
elemType := field.Type().Elem()
|
|
converted, err := convertValue(src, elemType)
|
|
if err != nil {
|
|
return err
|
|
}
|
|
ptr := reflect.New(elemType)
|
|
ptr.Elem().Set(converted)
|
|
field.Set(ptr)
|
|
return nil
|
|
}
|
|
converted, err := convertValue(src, field.Type())
|
|
if err == nil {
|
|
field.Set(converted)
|
|
return nil
|
|
}
|
|
if field.Type() == encryptedType {
|
|
return err
|
|
}
|
|
if field.CanAddr() {
|
|
if scanner, ok := field.Addr().Interface().(sql.Scanner); ok {
|
|
scanSrc, scanErr := fillScanSource(val)
|
|
if scanErr != nil {
|
|
return scanErr
|
|
}
|
|
return scanner.Scan(scanSrc)
|
|
}
|
|
}
|
|
return err
|
|
}
|
|
|
|
func fillScanSource(val any) (any, error) {
|
|
switch v := val.(type) {
|
|
case string, []byte:
|
|
return v, nil
|
|
default:
|
|
b, err := json.Marshal(v)
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
return b, nil
|
|
}
|
|
}
|
|
|
|
var encryptedType = reflect.TypeOf(Encrypted{})
|
|
|
|
func convertValue(src reflect.Value, destType reflect.Type) (reflect.Value, error) {
|
|
if destType == encryptedType {
|
|
return encryptedFromRequest(src)
|
|
}
|
|
if src.Type().AssignableTo(destType) {
|
|
return src, nil
|
|
}
|
|
if n, ok := src.Interface().(json.Number); ok {
|
|
if out, handled, err := convertNumber(n, destType); handled {
|
|
return out, err
|
|
}
|
|
}
|
|
if src.Type().ConvertibleTo(destType) {
|
|
return src.Convert(destType), nil
|
|
}
|
|
assignErr := fmt.Errorf("cannot assign %s to %s", src.Type(), destType)
|
|
if out, handled, err := convertText(src, destType); handled {
|
|
if err != nil {
|
|
return reflect.Value{}, fmt.Errorf("%w: %w", assignErr, err)
|
|
}
|
|
return out, nil
|
|
}
|
|
return reflect.Value{}, assignErr
|
|
}
|
|
|
|
var (
|
|
textUnmarshalerType = reflect.TypeOf((*encoding.TextUnmarshaler)(nil)).Elem()
|
|
scannerType = reflect.TypeOf((*sql.Scanner)(nil)).Elem()
|
|
dateType = reflect.TypeOf(Date{})
|
|
timeOfDayType = reflect.TypeOf(TimeOfDay{})
|
|
)
|
|
|
|
// convertText fills a destination type whose pointer implements
|
|
// encoding.TextUnmarshaler (time.Time, Date, TimeOfDay) from a string or
|
|
// []byte source, so a JSON date string fills a date column. It runs only
|
|
// after every other conversion failed, and it skips types that also
|
|
// implement sql.Scanner (other than Date and TimeOfDay), which setField
|
|
// already fills through Scan, so nothing that filled before changes path.
|
|
// handled is false when the fallback does not apply.
|
|
func convertText(src reflect.Value, destType reflect.Type) (reflect.Value, bool, error) {
|
|
var text []byte
|
|
switch v := src.Interface().(type) {
|
|
case string:
|
|
text = []byte(v)
|
|
case []byte:
|
|
text = v
|
|
default:
|
|
return reflect.Value{}, false, nil
|
|
}
|
|
ptrType := reflect.PointerTo(destType)
|
|
if !ptrType.Implements(textUnmarshalerType) {
|
|
return reflect.Value{}, false, nil
|
|
}
|
|
if ptrType.Implements(scannerType) && destType != dateType && destType != timeOfDayType {
|
|
return reflect.Value{}, false, nil
|
|
}
|
|
ptr := reflect.New(destType)
|
|
if err := ptr.Interface().(encoding.TextUnmarshaler).UnmarshalText(text); err != nil {
|
|
return reflect.Value{}, true, err
|
|
}
|
|
return ptr.Elem(), true, nil
|
|
}
|
|
|
|
// convertNumber parses a json.Number, which a request decoder using
|
|
// UseNumber produces, into an integer, unsigned or float field. A value that
|
|
// does not parse as that kind or overflows it is an error. handled is false
|
|
// for any other destination kind, so a string field still takes the number's
|
|
// text through the ordinary conversion.
|
|
func convertNumber(n json.Number, destType reflect.Type) (reflect.Value, bool, error) {
|
|
out := reflect.New(destType).Elem()
|
|
switch destType.Kind() {
|
|
case reflect.Int, reflect.Int8, reflect.Int16, reflect.Int32, reflect.Int64:
|
|
i, err := strconv.ParseInt(n.String(), 10, 64)
|
|
if err != nil || out.OverflowInt(i) {
|
|
return reflect.Value{}, true, fmt.Errorf("cannot assign number %s to %s", n, destType)
|
|
}
|
|
out.SetInt(i)
|
|
case reflect.Uint, reflect.Uint8, reflect.Uint16, reflect.Uint32, reflect.Uint64:
|
|
u, err := strconv.ParseUint(n.String(), 10, 64)
|
|
if err != nil || out.OverflowUint(u) {
|
|
return reflect.Value{}, true, fmt.Errorf("cannot assign number %s to %s", n, destType)
|
|
}
|
|
out.SetUint(u)
|
|
case reflect.Float32, reflect.Float64:
|
|
f, err := strconv.ParseFloat(n.String(), 64)
|
|
if err != nil || out.OverflowFloat(f) {
|
|
return reflect.Value{}, true, fmt.Errorf("cannot assign number %s to %s", n, destType)
|
|
}
|
|
out.SetFloat(f)
|
|
default:
|
|
return reflect.Value{}, false, nil
|
|
}
|
|
return out, true, nil
|
|
}
|
|
|
|
// encryptedFromRequest treats request input for an Encrypted column as
|
|
// plaintext. It never falls through to Encrypted.Scan: Scan decrypts, so a
|
|
// write path that scanned request input would reject real secrets and accept
|
|
// another row's ciphertext, copying that row's secret.
|
|
func encryptedFromRequest(src reflect.Value) (reflect.Value, error) {
|
|
if src.Type() == encryptedType {
|
|
return src, nil
|
|
}
|
|
if src.Kind() != reflect.String {
|
|
return reflect.Value{}, fmt.Errorf("encrypted value must be a string")
|
|
}
|
|
return reflect.ValueOf(NewEncrypted(src.String())), nil
|
|
}
|