@@ 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 @@ func (rid *RID) UnmarshalGQL(v any) error {
// 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