package api_gen import ( "fmt" "go/ast" "go/constant" "go/token" "go/types" "strconv" "strings" "golang.org/x/tools/go/packages" ) // FilterSchema is the parsed representation of a filter.Schema. type FilterSchema struct { DefaultFields []string `json:"DefaultFields,omitempty"` Fields []FilterField `json:"Fields"` } // FilterField is the parsed representation of a filter.Field. type FilterField struct { Name string `json:"Name"` Type string `json:"Type"` Values []string `json:"Values,omitempty"` } // collectFilterSchemas finds all vars named *Filter in pkg by naming convention, // parses their filter.Schema values from AST, and returns a map of endpoint name → FilterSchema. func collectFilterSchemas(pkg *packages.Package) (map[string]FilterSchema, error) { allPkgs := collectAllPkgs(pkg) schemas := map[string]FilterSchema{} scope := pkg.Types.Scope() for _, name := range scope.Names() { if !strings.HasSuffix(name, "Filter") { continue } obj := scope.Lookup(name) if _, ok := obj.(*types.Var); !ok { continue } endpointName := strings.TrimSuffix(name, "Filter") schema, err := parseFilterSchemaVar(pkg, allPkgs, name) if err != nil { return nil, fmt.Errorf("parse %s: %w", name, err) } schemas[endpointName] = *schema } return schemas, nil } func collectAllPkgs(pkg *packages.Package) map[string]*packages.Package { all := map[string]*packages.Package{} var collect func(p *packages.Package) collect = func(p *packages.Package) { if _, seen := all[p.PkgPath]; seen { return } all[p.PkgPath] = p for _, imp := range p.Imports { collect(imp) } } collect(pkg) return all } func parseFilterSchemaVar(pkg *packages.Package, allPkgs map[string]*packages.Package, varName string) (*FilterSchema, error) { initExpr, err := findVarInitExpr(pkg, varName) if err != nil { return nil, err } return evalFilterSchema(initExpr, pkg, allPkgs) } func findVarInitExpr(pkg *packages.Package, varName string) (ast.Expr, error) { for _, file := range pkg.Syntax { for _, decl := range file.Decls { gd, ok := decl.(*ast.GenDecl) if !ok || gd.Tok != token.VAR { continue } for _, spec := range gd.Specs { vs, ok := spec.(*ast.ValueSpec) if !ok { continue } for i, ident := range vs.Names { if ident.Name == varName && i < len(vs.Values) { return vs.Values[i], nil } } } } } return nil, fmt.Errorf("var %s not found in syntax", varName) } // evalFilterSchema evaluates an expression to a FilterSchema. It handles // composite literals, selector expressions (cross-package var references), // and builder-style method chains (Schema{}.String(...).Enum(...).Default(...)). func evalFilterSchema(expr ast.Expr, pkg *packages.Package, allPkgs map[string]*packages.Package) (*FilterSchema, error) { switch e := expr.(type) { case *ast.CompositeLit: return parseFilterSchemaLit(e, pkg) case *ast.SelectorExpr: obj := pkg.TypesInfo.Uses[e.Sel] if obj == nil { return nil, fmt.Errorf("could not resolve %s", e.Sel.Name) } targetPkg, ok := allPkgs[obj.Pkg().Path()] if !ok { return nil, fmt.Errorf("package %s not found", obj.Pkg().Path()) } targetExpr, err := findVarInitExpr(targetPkg, e.Sel.Name) if err != nil { return nil, err } return evalFilterSchema(targetExpr, targetPkg, allPkgs) case *ast.CallExpr: return evalFilterSchemaChain(e, pkg, allPkgs) } return nil, fmt.Errorf("unexpected expression type %T", expr) } // evalFilterSchemaChain evaluates one step of a Schema builder chain and // recurses into the receiver, building up the FilterSchema left-to-right. func evalFilterSchemaChain(call *ast.CallExpr, pkg *packages.Package, allPkgs map[string]*packages.Package) (*FilterSchema, error) { sel, ok := call.Fun.(*ast.SelectorExpr) if !ok { return nil, fmt.Errorf("filter schema chain: expected selector, got %T", call.Fun) } schema, err := evalFilterSchema(sel.X, pkg, allPkgs) if err != nil { return nil, err } switch sel.Sel.Name { case "String": if len(call.Args) != 1 { return nil, fmt.Errorf("Schema.String expects 1 argument") } name, err := evalStringExpr(call.Args[0], pkg.TypesInfo) if err != nil { return nil, fmt.Errorf("Schema.String name: %w", err) } schema.Fields = append(schema.Fields, FilterField{Name: name, Type: "string"}) case "Enum": if len(call.Args) != 2 { return nil, fmt.Errorf("Schema.Enum expects 2 arguments") } name, err := evalStringExpr(call.Args[0], pkg.TypesInfo) if err != nil { return nil, fmt.Errorf("Schema.Enum name: %w", err) } values, err := evalEnumValues(call.Args[1], pkg, allPkgs) if err != nil { return nil, fmt.Errorf("Schema.Enum values: %w", err) } schema.Fields = append(schema.Fields, FilterField{Name: name, Type: "enum", Values: values}) case "Default": if len(call.Args) != 1 { return nil, fmt.Errorf("Schema.Default expects 1 argument") } name, err := evalStringExpr(call.Args[0], pkg.TypesInfo) if err != nil { return nil, fmt.Errorf("Schema.Default name: %w", err) } schema.DefaultFields = append(schema.DefaultFields, name) default: return nil, fmt.Errorf("unknown filter.Schema method %q", sel.Sel.Name) } return schema, nil } // evalEnumValues extracts the valid string values from an atlas9 enum type. // expr must be a value of the enum type (e.g. StatusUnknown). It finds the // generated allXxx variable in the enum's package and reads the string // literals from each non-zero element. func evalEnumValues(expr ast.Expr, pkg *packages.Package, allPkgs map[string]*packages.Package) ([]string, error) { var named *types.Named switch e := expr.(type) { case *ast.Ident: if obj := pkg.TypesInfo.Uses[e]; obj != nil { named, _ = obj.Type().(*types.Named) } case *ast.SelectorExpr: if obj := pkg.TypesInfo.Uses[e.Sel]; obj != nil { named, _ = obj.Type().(*types.Named) } } if named == nil { return nil, fmt.Errorf("cannot determine enum type from %T", expr) } typeName := named.Obj().Name() pkgPath := named.Obj().Pkg().Path() targetPkg, ok := allPkgs[pkgPath] if !ok { return nil, fmt.Errorf("package %s not found", pkgPath) } // The enumgen template generates: var allStatus = []Status{StatusUnknown, StatusOpen, ...} allVarName := "all" + typeName allExpr, err := findVarInitExpr(targetPkg, allVarName) if err != nil { return nil, fmt.Errorf("finding %s: %w", allVarName, err) } allLit, ok := allExpr.(*ast.CompositeLit) if !ok { return nil, fmt.Errorf("%s is not a composite literal", allVarName) } var values []string for i, elt := range allLit.Elts { if i == 0 { continue // skip the zero/Unknown value } eltLit, ok := elt.(*ast.CompositeLit) if !ok || len(eltLit.Elts) == 0 { continue } s, err := evalStringExpr(eltLit.Elts[0], targetPkg.TypesInfo) if err != nil { continue } values = append(values, s) } return values, nil } func parseFilterSchemaLit(lit *ast.CompositeLit, pkg *packages.Package) (*FilterSchema, error) { schema := &FilterSchema{} for _, elt := range lit.Elts { kv, ok := elt.(*ast.KeyValueExpr) if !ok { continue } key, ok := kv.Key.(*ast.Ident) if !ok { continue } switch key.Name { case "DefaultFields": vals, err := evalStringSlice(kv.Value, pkg.TypesInfo) if err != nil { return nil, fmt.Errorf("DefaultFields: %w", err) } schema.DefaultFields = vals case "Fields": fields, err := parseFilterFieldSlice(kv.Value, pkg.TypesInfo) if err != nil { return nil, fmt.Errorf("Fields: %w", err) } schema.Fields = fields } } return schema, nil } func parseFilterFieldSlice(expr ast.Expr, info *types.Info) ([]FilterField, error) { lit, ok := expr.(*ast.CompositeLit) if !ok { return nil, fmt.Errorf("expected composite literal, got %T", expr) } var fields []FilterField for _, elt := range lit.Elts { fieldLit, ok := elt.(*ast.CompositeLit) if !ok { return nil, fmt.Errorf("expected field literal, got %T", elt) } f, err := parseFilterFieldLit(fieldLit, info) if err != nil { return nil, err } fields = append(fields, f) } return fields, nil } func parseFilterFieldLit(lit *ast.CompositeLit, info *types.Info) (FilterField, error) { f := FilterField{Type: "string"} for _, elt := range lit.Elts { kv, ok := elt.(*ast.KeyValueExpr) if !ok { continue } key, ok := kv.Key.(*ast.Ident) if !ok { continue } switch key.Name { case "Name": s, err := evalStringExpr(kv.Value, info) if err != nil { return f, fmt.Errorf("Name: %w", err) } f.Name = s case "Type": t, err := evalFieldType(kv.Value, info) if err != nil { return f, fmt.Errorf("Type: %w", err) } f.Type = t case "Values": vals, err := evalStringSlice(kv.Value, info) if err != nil { return f, fmt.Errorf("Values: %w", err) } f.Values = vals } } return f, nil } func evalStringSlice(expr ast.Expr, info *types.Info) ([]string, error) { lit, ok := expr.(*ast.CompositeLit) if !ok { return nil, fmt.Errorf("expected composite literal, got %T", expr) } var result []string for _, elt := range lit.Elts { s, err := evalStringExpr(elt, info) if err != nil { return nil, err } result = append(result, s) } return result, nil } func evalStringExpr(expr ast.Expr, info *types.Info) (string, error) { switch e := expr.(type) { case *ast.BasicLit: if e.Kind == token.STRING { return strconv.Unquote(e.Value) } case *ast.CallExpr: // string(X) conversion — evaluate the argument as a constant if ident, ok := e.Fun.(*ast.Ident); ok && ident.Name == "string" && len(e.Args) == 1 { return evalStringExpr(e.Args[0], info) } case *ast.Ident: if obj := info.Uses[e]; obj != nil { if c, ok := obj.(*types.Const); ok { return constant.StringVal(c.Val()), nil } } case *ast.SelectorExpr: if obj := info.Uses[e.Sel]; obj != nil { if c, ok := obj.(*types.Const); ok { return constant.StringVal(c.Val()), nil } } } return "", fmt.Errorf("cannot evaluate as string: %T", expr) } func evalFieldType(expr ast.Expr, info *types.Info) (string, error) { var obj types.Object switch e := expr.(type) { case *ast.Ident: obj = info.Uses[e] case *ast.SelectorExpr: obj = info.Uses[e.Sel] } if obj == nil { return "", fmt.Errorf("cannot resolve field type from %T", expr) } c, ok := obj.(*types.Const) if !ok { return "", fmt.Errorf("field type is not a constant") } v, exact := constant.Int64Val(c.Val()) if !exact { return "", fmt.Errorf("field type constant is not an integer") } switch v { case 0: return "string", nil case 1: return "enum", nil } return "", fmt.Errorf("unknown field type value %d", v) }