mirror of
https://github.com/coder/coder.git
synced 2026-09-24 15:04:27 +08:00
chore: add custom querier functions to dbgen (#8496)
* chore: add custom querier functions to dbgen * chore: parse package was missing some imports, so force them
This commit is contained in:
+58
-7
@@ -418,21 +418,44 @@ type querierFunction struct {
|
||||
|
||||
// readQuerierFunctions reads the functions from coderd/database/querier.go
|
||||
func readQuerierFunctions() ([]querierFunction, error) {
|
||||
f, err := parseDBFile("querier.go")
|
||||
if err != nil {
|
||||
return nil, xerrors.Errorf("parse querier.go: %w", err)
|
||||
}
|
||||
funcs, err := loadInterfaceFuncs(f, "sqlcQuerier")
|
||||
if err != nil {
|
||||
return nil, xerrors.Errorf("load interface %s funcs: %w", "sqlcQuerier", err)
|
||||
}
|
||||
|
||||
customFile, err := parseDBFile("modelqueries.go")
|
||||
if err != nil {
|
||||
return nil, xerrors.Errorf("parse modelqueriers.go: %w", err)
|
||||
}
|
||||
// Custom funcs should be appended after the regular functions
|
||||
customFuncs, err := loadInterfaceFuncs(customFile, "customQuerier")
|
||||
if err != nil {
|
||||
return nil, xerrors.Errorf("load interface %s funcs: %w", "customQuerier", err)
|
||||
}
|
||||
|
||||
return append(funcs, customFuncs...), nil
|
||||
}
|
||||
|
||||
func parseDBFile(filename string) (*dst.File, error) {
|
||||
localPath, err := localFilePath()
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
querierPath := filepath.Join(localPath, "..", "..", "..", "coderd", "database", "querier.go")
|
||||
|
||||
querierPath := filepath.Join(localPath, "..", "..", "..", "coderd", "database", filename)
|
||||
querierData, err := os.ReadFile(querierPath)
|
||||
if err != nil {
|
||||
return nil, xerrors.Errorf("read querier: %w", err)
|
||||
return nil, xerrors.Errorf("read %s: %w", filename, err)
|
||||
}
|
||||
f, err := decorator.Parse(querierData)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return f, err
|
||||
}
|
||||
|
||||
func loadInterfaceFuncs(f *dst.File, interfaceName string) ([]querierFunction, error) {
|
||||
var querier *dst.InterfaceType
|
||||
for _, decl := range f.Decls {
|
||||
genDecl, ok := decl.(*dst.GenDecl)
|
||||
@@ -447,7 +470,7 @@ func readQuerierFunctions() ([]querierFunction, error) {
|
||||
}
|
||||
// This is the name of the interface. If that ever changes,
|
||||
// this will need to be updated.
|
||||
if typeSpec.Name.Name != "sqlcQuerier" {
|
||||
if typeSpec.Name.Name != interfaceName {
|
||||
continue
|
||||
}
|
||||
querier, ok = typeSpec.Type.(*dst.InterfaceType)
|
||||
@@ -461,7 +484,8 @@ func readQuerierFunctions() ([]querierFunction, error) {
|
||||
return nil, xerrors.Errorf("querier not found")
|
||||
}
|
||||
funcs := []querierFunction{}
|
||||
for _, method := range querier.Methods.List {
|
||||
allMethods := interfaceMethods(querier)
|
||||
for _, method := range allMethods {
|
||||
funcType, ok := method.Type.(*dst.FuncType)
|
||||
if !ok {
|
||||
continue
|
||||
@@ -540,3 +564,30 @@ func nameFromSnakeCase(s string) string {
|
||||
}
|
||||
return ret
|
||||
}
|
||||
|
||||
// interfaceMethods returns all embedded methods of an interface.
|
||||
func interfaceMethods(i *dst.InterfaceType) []*dst.Field {
|
||||
var allMethods []*dst.Field
|
||||
for _, field := range i.Methods.List {
|
||||
switch fieldType := field.Type.(type) {
|
||||
case *dst.FuncType:
|
||||
allMethods = append(allMethods, field)
|
||||
case *dst.InterfaceType:
|
||||
allMethods = append(allMethods, interfaceMethods(fieldType)...)
|
||||
case *dst.Ident:
|
||||
// Embedded interfaces are Idents -> TypeSpec -> InterfaceType
|
||||
// If the embedded interface is not in the parsed file, then
|
||||
// the Obj will be nil.
|
||||
if fieldType.Obj != nil {
|
||||
objDecl, ok := fieldType.Obj.Decl.(*dst.TypeSpec)
|
||||
if ok {
|
||||
isInterface, ok := objDecl.Type.(*dst.InterfaceType)
|
||||
if ok {
|
||||
allMethods = append(allMethods, interfaceMethods(isInterface)...)
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
return allMethods
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user