diff --git a/model/rid.go b/model/rid.go index 15b12401c1b25318da82093d7b27033ff9665874..8af5a9d67e517b9ac8605f23aaec173c3e38bacc 100644 --- a/model/rid.go +++ b/model/rid.go @@ -44,18 +44,22 @@ func (rid RID) MarshalGQL(w io.Writer) { w.Write(fmt.Appendf(nil, `"%s"`, rid.String())) } +func (rid *RID) unmarshalString(s string) error { + bytes, err := base32Encoding.DecodeString(s) + if err != nil { + return err + } + rid.uuid, err = uuid.FromBytes(bytes) + if err != nil { + return err + } + return nil +} + func (rid *RID) UnmarshalGQL(v any) error { switch v := v.(type) { case string: - bytes, err := base32Encoding.DecodeString(v) - if err != nil { - return err - } - rid.uuid, err = uuid.FromBytes(bytes) - if err != nil { - return err - } - return nil + return rid.unmarshalString(v) default: return fmt.Errorf("%T is not a valid RID", v) } @@ -63,13 +67,18 @@ } // database/sql.Scanner func (rid *RID) Scan(src any) error { - var uuid uuid.UUID - err := uuid.Scan(src) - if err != nil { - return err + switch src := src.(type) { + case string: + return rid.unmarshalString(src) + default: + var uuid uuid.UUID + err := uuid.Scan(src) + if err != nil { + return err + } + rid.uuid = uuid + return nil } - rid.uuid = uuid - return nil } // database/sql/driver.Valuer