package database import ( "context" "sort" "strings" "github.com/lib/pq" "github.com/vektah/gqlparser/v2/ast" "github.com/99designs/gqlgen/graphql" ) func collectFields(ctx context.Context) []graphql.CollectedField { var fields []graphql.CollectedField if graphql.GetFieldContext(ctx) != nil { fields = graphql.CollectFieldsCtx(ctx, nil) octx := graphql.GetOperationContext(ctx) for _, col := range fields { if col.Name == "results" { // This endpoint is using the cursor pattern; the columns we // actually need to filter with are nested into the results // field. fields = graphql.CollectFields(octx, col.SelectionSet, nil) break } } } return fields } func Scan(ctx context.Context, m Model) []interface{} { qlFields := collectFields(ctx) if len(qlFields) == 0 { // Collect all fields if we are not in an active graphql context for _, field := range m.Fields().All() { qlFields = append(qlFields, graphql.CollectedField{ Field: &ast.Field{Name: field.GQL}, }) } } sort.Slice(qlFields, func(a, b int) bool { return qlFields[a].Name < qlFields[b].Name }) var fields []interface{} for _, qlField := range qlFields { if gqlFields, ok := m.Fields().GQL(qlField.Name); ok { for _, field := range gqlFields { fields = append(fields, field.Ptr) } } } for _, field := range m.Fields().Anonymous() { fields = append(fields, field.Ptr) } return fields } func Columns(ctx context.Context, m Model) []string { fields := collectFields(ctx) if len(fields) == 0 { // Collect all fields if we are not in an active graphql context for _, field := range m.Fields().All() { fields = append(fields, graphql.CollectedField{ Field: &ast.Field{Name: field.GQL}, }) } } sort.Slice(fields, func(a, b int) bool { return fields[a].Name < fields[b].Name }) var columns []string for _, gql := range fields { if sqlFields, ok := m.Fields().GQL(gql.Name); ok { for _, field := range sqlFields { sql := field.SQL if !strings.ContainsRune(sql, '.') { sql = WithAlias(m.Alias(), field.SQL) } columns = append(columns, sql) } } } for _, field := range m.Fields().Anonymous() { alias := m.Alias() if alias == "" { alias = m.Table() } columns = append(columns, WithAlias(alias, field.SQL)) } return columns } func WithAlias(alias, col string) string { if alias != "" { return pq.QuoteIdentifier(alias) + "." + pq.QuoteIdentifier(col) } else { return pq.QuoteIdentifier(col) } }