@@ -185,6 +185,7 @@ func argName(name string) string {
185185func buildQueries (req * plugin.GenerateRequest , options * opts.Options , enums []Enum , structs []Struct ) ([]Query , error ) {
186186 models := buildModelTypeSet (enums , structs )
187187 qualifier := options .ModelsTypeQualifier ()
188+ rowTypes := newRowTypes (structs )
188189 qs := make ([]Query , 0 , len (req .Queries ))
189190 for _ , query := range req .Queries {
190191 if query .Name == "" {
@@ -269,6 +270,10 @@ func buildQueries(req *plugin.GenerateRequest, options *opts.Options, enums []En
269270 }
270271 }
271272
273+ if query .TypeName != "" && ! returnsStruct (query ) {
274+ return nil , fmt .Errorf ("query %s: :type %s needs a query that returns more than one column" , query .Name , query .TypeName )
275+ }
276+
272277 if len (query .Columns ) == 1 && query .Columns [0 ].EmbedTable == nil {
273278 c := query .Columns [0 ]
274279 name := columnName (c , 0 )
@@ -303,7 +308,7 @@ func buildQueries(req *plugin.GenerateRequest, options *opts.Options, enums []En
303308 var emit bool
304309
305310 for _ , s := range structs {
306- if len (s .Fields ) != len (query .Columns ) {
311+ if query . TypeName != "" || len (s .Fields ) != len (query .Columns ) {
307312 continue
308313 }
309314 same := true
@@ -331,12 +336,22 @@ func buildQueries(req *plugin.GenerateRequest, options *opts.Options, enums []En
331336 embed : newGoEmbed (c .EmbedTable , structs , req .Catalog .DefaultSchema ),
332337 })
333338 }
339+ name := gq .MethodName + "Row"
340+ if query .TypeName != "" {
341+ name = query .TypeName
342+ }
334343 var err error
335- gs , err = columnsToStruct (req , options , gq . MethodName + "Row" , columns , true , models , qualifier )
344+ gs , err = columnsToStruct (req , options , name , columns , true , models , qualifier )
336345 if err != nil {
337346 return nil , err
338347 }
339348 emit = true
349+ if query .TypeName != "" {
350+ gs , emit , err = rowTypes .add (query .Name , gs )
351+ if err != nil {
352+ return nil , err
353+ }
354+ }
340355 }
341356 gq .Ret = QueryValue {
342357 Emit : emit ,
@@ -361,6 +376,76 @@ var cmdReturnsData = map[string]struct{}{
361376 metadata .CmdOne : {},
362377}
363378
379+ // returnsStruct reports whether the Go code for query returns its rows as a
380+ // struct, which is what a ":type" annotation names.
381+ func returnsStruct (query * plugin.Query ) bool {
382+ if len (query .Columns ) == 1 && query .Columns [0 ].EmbedTable == nil {
383+ return false
384+ }
385+ return putOutColumns (query )
386+ }
387+
388+ // rowTypes tracks the structs named by ":type" annotations, so that the
389+ // queries sharing a name return one struct.
390+ type rowTypes struct {
391+ models map [string ]bool
392+ named map [string ]rowType
393+ }
394+
395+ type rowType struct {
396+ query string // the first query to use the name, which emits the struct
397+ s * Struct
398+ }
399+
400+ func newRowTypes (models []Struct ) * rowTypes {
401+ r := & rowTypes {models : map [string ]bool {}, named : map [string ]rowType {}}
402+ for _ , m := range models {
403+ r .models [m .Name ] = true
404+ }
405+ return r
406+ }
407+
408+ // add takes the struct built for the rows of a query annotated with ":type"
409+ // and returns the struct the query returns and whether to emit it. The first
410+ // query to use a name emits its struct; the others return that struct, as long
411+ // as their columns give the same fields.
412+ func (r * rowTypes ) add (query string , s * Struct ) (* Struct , bool , error ) {
413+ if r .models [s .Name ] {
414+ return nil , false , fmt .Errorf ("query %s: :type %s is already the name of a model" , query , s .Name )
415+ }
416+ first , ok := r .named [s .Name ]
417+ if ! ok {
418+ r .named [s .Name ] = rowType {query : query , s : s }
419+ return s , true , nil
420+ }
421+ if diff := fieldsDiff (first .s .Fields , s .Fields ); diff != "" {
422+ return nil , false , fmt .Errorf ("query %s: :type %s does not match query %s: %s" , query , s .Name , first .query , diff )
423+ }
424+ return first .s , false , nil
425+ }
426+
427+ // fieldsDiff describes the first difference of got from want, or returns ""
428+ // if they are the same.
429+ func fieldsDiff (want , got []Field ) string {
430+ if len (got ) != len (want ) {
431+ return fmt .Sprintf ("%d columns, want %d" , len (got ), len (want ))
432+ }
433+ for i , w := range want {
434+ g := got [i ]
435+ if g .Name != w .Name || g .Type != w .Type || g .Tag () != w .Tag () || fieldsDiff (w .EmbedFields , g .EmbedFields ) != "" {
436+ return fmt .Sprintf ("column %d is %s, want %s" , i + 1 , describeField (g ), describeField (w ))
437+ }
438+ }
439+ return ""
440+ }
441+
442+ func describeField (f Field ) string {
443+ if tag := f .Tag (); tag != "" {
444+ return fmt .Sprintf ("%s %s `%s`" , f .Name , f .Type , tag )
445+ }
446+ return f .Name + " " + f .Type
447+ }
448+
364449func putOutColumns (query * plugin.Query ) bool {
365450 _ , found := cmdReturnsData [query .Cmd ]
366451 return found
0 commit comments