1
0

walker.go 5.6 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249
  1. package main
  2. import (
  3. "fmt"
  4. "go/ast"
  5. "go/parser"
  6. "go/token"
  7. "io/fs"
  8. "path/filepath"
  9. "strings"
  10. )
  11. type walkOverride struct {
  12. Field string
  13. Kind TypeKind
  14. }
  15. type packageRequest struct {
  16. Path string
  17. StructAllow map[string]bool
  18. AliasAllow map[string]bool
  19. Overrides map[string][]walkOverride
  20. }
  21. func walkPackages(requests []packageRequest) ([]Schema, []Alias, error) {
  22. fset := token.NewFileSet()
  23. var schemas []Schema
  24. var aliases []Alias
  25. for _, req := range requests {
  26. dir := req.Path
  27. pkgs, err := parser.ParseDir(fset, dir, func(fi fs.FileInfo) bool {
  28. return !strings.HasSuffix(fi.Name(), "_test.go")
  29. }, parser.ParseComments)
  30. if err != nil {
  31. return nil, nil, fmt.Errorf("parse %s: %w", dir, err)
  32. }
  33. for _, pkg := range pkgs {
  34. for _, file := range pkg.Files {
  35. for _, decl := range file.Decls {
  36. gen, ok := decl.(*ast.GenDecl)
  37. if !ok || gen.Tok != token.TYPE {
  38. continue
  39. }
  40. for _, spec := range gen.Specs {
  41. ts, ok := spec.(*ast.TypeSpec)
  42. if !ok {
  43. continue
  44. }
  45. if strct, ok := ts.Type.(*ast.StructType); ok {
  46. if req.StructAllow != nil && !req.StructAllow[ts.Name.Name] {
  47. continue
  48. }
  49. s := Schema{
  50. Name: ts.Name.Name,
  51. Package: pkg.Name,
  52. Doc: collectDoc(gen.Doc, ts.Doc),
  53. }
  54. overrides := req.Overrides[ts.Name.Name]
  55. for _, fld := range strct.Fields.List {
  56. s.Fields = append(s.Fields, buildFields(fld, overrides)...)
  57. }
  58. schemas = append(schemas, s)
  59. continue
  60. }
  61. if _, ok := ts.Type.(*ast.InterfaceType); ok {
  62. continue
  63. }
  64. if req.AliasAllow != nil && !req.AliasAllow[ts.Name.Name] {
  65. continue
  66. }
  67. aliases = append(aliases, Alias{
  68. Name: ts.Name.Name,
  69. Package: pkg.Name,
  70. Underlying: exprToType(ts.Type),
  71. })
  72. }
  73. }
  74. }
  75. }
  76. }
  77. return schemas, aliases, nil
  78. }
  79. func collectDoc(group ...*ast.CommentGroup) string {
  80. var b strings.Builder
  81. for _, g := range group {
  82. if g == nil {
  83. continue
  84. }
  85. for _, c := range g.List {
  86. line := strings.TrimPrefix(c.Text, "// ")
  87. line = strings.TrimPrefix(line, "//")
  88. b.WriteString(strings.TrimSpace(line))
  89. b.WriteByte('\n')
  90. }
  91. }
  92. return strings.TrimSpace(b.String())
  93. }
  94. func buildFields(fld *ast.Field, overrides []walkOverride) []Field {
  95. var fields []Field
  96. tag := ""
  97. if fld.Tag != nil {
  98. tag = fld.Tag.Value
  99. }
  100. jsonTag, validateTag, exampleTag, gormDash := parseStructTag(tag)
  101. if gormDash && jsonTag == "" {
  102. return nil
  103. }
  104. jsonName, omit, omitempty := parseJSONTag(jsonTag)
  105. if omit {
  106. return nil
  107. }
  108. validate := parseValidateTag(validateTag)
  109. doc := collectDoc(fld.Doc, fld.Comment)
  110. for _, n := range fld.Names {
  111. fname := jsonName
  112. if fname == "" {
  113. fname = lowerFirst(n.Name)
  114. }
  115. t := exprToType(fld.Type)
  116. for _, o := range overrides {
  117. if o.Field == n.Name || o.Field == jsonName {
  118. t = TypeRef{Kind: o.Kind}
  119. break
  120. }
  121. }
  122. fields = append(fields, Field{
  123. JSONName: fname,
  124. GoName: n.Name,
  125. Type: t,
  126. Optional: omitempty || isPointer(fld.Type),
  127. Validate: validate,
  128. Doc: doc,
  129. Example: exampleTag,
  130. })
  131. }
  132. if len(fld.Names) == 0 {
  133. fname := jsonName
  134. if fname == "" {
  135. fname = lowerFirst(exprIdentName(fld.Type))
  136. }
  137. t := exprToType(fld.Type)
  138. for _, o := range overrides {
  139. if o.Field == exprIdentName(fld.Type) || o.Field == jsonName {
  140. t = TypeRef{Kind: o.Kind}
  141. break
  142. }
  143. }
  144. fields = append(fields, Field{
  145. JSONName: fname,
  146. GoName: exprIdentName(fld.Type),
  147. Type: t,
  148. Optional: omitempty || isPointer(fld.Type),
  149. Validate: validate,
  150. Doc: doc,
  151. Example: exampleTag,
  152. })
  153. }
  154. return fields
  155. }
  156. func exprToType(expr ast.Expr) TypeRef {
  157. switch e := expr.(type) {
  158. case *ast.Ident:
  159. return identType(e.Name)
  160. case *ast.StarExpr:
  161. inner := exprToType(e.X)
  162. return TypeRef{Kind: KindRef, Name: "nullable", Inner: &inner}
  163. case *ast.ArrayType:
  164. elem := exprToType(e.Elt)
  165. return TypeRef{Kind: KindArray, Element: &elem}
  166. case *ast.MapType:
  167. k := exprToType(e.Key)
  168. v := exprToType(e.Value)
  169. return TypeRef{Kind: KindMap, Key: &k, Value: &v}
  170. case *ast.SelectorExpr:
  171. pkg := exprIdentName(e.X)
  172. name := e.Sel.Name
  173. if pkg == "json" && name == "RawMessage" {
  174. return TypeRef{Kind: KindAny}
  175. }
  176. if pkg == "time" && name == "Time" {
  177. return TypeRef{Kind: KindString, Name: "datetime"}
  178. }
  179. return TypeRef{Kind: KindRef, Name: name}
  180. case *ast.InterfaceType:
  181. return TypeRef{Kind: KindAny}
  182. default:
  183. return TypeRef{Kind: KindUnknown}
  184. }
  185. }
  186. func identType(name string) TypeRef {
  187. switch name {
  188. case "string":
  189. return TypeRef{Kind: KindString}
  190. case "bool":
  191. return TypeRef{Kind: KindBool}
  192. case "int64", "uint64":
  193. return TypeRef{Kind: KindInt, Name: "int64"}
  194. case "int", "int8", "int16", "int32",
  195. "uint", "uint8", "uint16", "uint32":
  196. return TypeRef{Kind: KindInt}
  197. case "float32", "float64":
  198. return TypeRef{Kind: KindNumber}
  199. case "byte", "rune":
  200. return TypeRef{Kind: KindInt}
  201. case "any":
  202. return TypeRef{Kind: KindAny}
  203. default:
  204. return TypeRef{Kind: KindRef, Name: name}
  205. }
  206. }
  207. func isPointer(expr ast.Expr) bool {
  208. _, ok := expr.(*ast.StarExpr)
  209. return ok
  210. }
  211. func exprIdentName(expr ast.Expr) string {
  212. switch e := expr.(type) {
  213. case *ast.Ident:
  214. return e.Name
  215. case *ast.SelectorExpr:
  216. return e.Sel.Name
  217. case *ast.StarExpr:
  218. return exprIdentName(e.X)
  219. default:
  220. return ""
  221. }
  222. }
  223. func lowerFirst(s string) string {
  224. if s == "" {
  225. return s
  226. }
  227. return strings.ToLower(s[:1]) + s[1:]
  228. }
  229. func resolveRel(base, rel string) string {
  230. if filepath.IsAbs(rel) {
  231. return rel
  232. }
  233. return filepath.Clean(filepath.Join(base, rel))
  234. }