310 lines
8.7 KiB
Go
310 lines
8.7 KiB
Go
package main
|
||
|
||
import (
|
||
"bytes"
|
||
"errors"
|
||
"fmt"
|
||
"go/format"
|
||
"io"
|
||
"os"
|
||
"path/filepath"
|
||
"regexp"
|
||
"strings"
|
||
"unicode"
|
||
|
||
"git.apinb.com/bsm-tools/protoc-gen-slc/tpl"
|
||
|
||
"git.apinb.com/bsm-sdk/core/utils"
|
||
"golang.org/x/mod/modfile"
|
||
"google.golang.org/protobuf/compiler/protogen"
|
||
"google.golang.org/protobuf/types/pluginpb"
|
||
)
|
||
|
||
var ServicesName []string
|
||
|
||
func main() {
|
||
protogen.Options{}.Run(func(gen *protogen.Plugin) error {
|
||
gen.SupportedFeatures = uint64(pluginpb.CodeGeneratorResponse_FEATURE_PROTO3_OPTIONAL)
|
||
|
||
// 以代码生成根目录为基准;gen.Files 无 service 时不创建 internal/{server,logic}
|
||
serviceCount := 0
|
||
for _, f := range gen.Files {
|
||
serviceCount += len(f.Services)
|
||
}
|
||
if serviceCount == 0 {
|
||
return nil
|
||
}
|
||
|
||
if !utils.PathExists("./internal") {
|
||
os.MkdirAll("./internal", 0755)
|
||
}
|
||
if !utils.PathExists("./internal/server") {
|
||
os.MkdirAll("./internal/server", 0755)
|
||
}
|
||
if !utils.PathExists("./internal/logic") {
|
||
os.MkdirAll("./internal/logic", 0755)
|
||
}
|
||
|
||
for _, f := range gen.Files {
|
||
if len(f.Services) == 0 {
|
||
continue
|
||
}
|
||
if err := generateFiles(gen, f); err != nil {
|
||
return err
|
||
}
|
||
}
|
||
|
||
return generateNewServerFile(ServicesName)
|
||
})
|
||
}
|
||
|
||
func generateFiles(gen *protogen.Plugin, file *protogen.File) error {
|
||
for _, service := range file.Services {
|
||
ServicesName = append(ServicesName, service.GoName)
|
||
// Generate server file
|
||
if err := generateServerFile(gen, file, service); err != nil {
|
||
return err
|
||
}
|
||
|
||
// Generate logic file
|
||
if err := generateLogicFile(gen, file, service); err != nil {
|
||
return err
|
||
}
|
||
}
|
||
|
||
return nil
|
||
}
|
||
|
||
func generateNewServerFile(services []string) error {
|
||
moduleName := getModuleName()
|
||
|
||
//create new.go
|
||
code := tpl.NewFile
|
||
newImports := []string{
|
||
"pb \"" + moduleName + "/pb\"",
|
||
}
|
||
code = strings.ReplaceAll(code, "{import}", strings.Join(newImports, "\n"))
|
||
|
||
// register grpc
|
||
var register []string
|
||
for _, service := range services {
|
||
register = append(register, "pb.Register"+service+"Server(srv.Grpc, New"+service+"Server())")
|
||
}
|
||
code = strings.ReplaceAll(code, "{register}", strings.Join(register, "\n"))
|
||
|
||
// register grpc gw
|
||
var gw []string
|
||
for _, service := range services {
|
||
handler := strings.ReplaceAll(tpl.Handler, "{service}", service)
|
||
gw = append(gw, handler)
|
||
}
|
||
code = strings.ReplaceAll(code, "{gw}", strings.Join(gw, "\n"))
|
||
|
||
// 格式化代码
|
||
formattedCode, err := format.Source([]byte(code))
|
||
if err != nil {
|
||
return fmt.Errorf("failed to format generated code: %w", err)
|
||
}
|
||
|
||
StringToFile("./internal/server/new.go", string(formattedCode))
|
||
|
||
return nil
|
||
}
|
||
|
||
func generateServerFile(gen *protogen.Plugin, file *protogen.File, service *protogen.Service) error {
|
||
filename := fmt.Sprintf("./internal/server/%s_server.go", toSnakeCase(service.GoName))
|
||
moduleName := getModuleName()
|
||
|
||
//create servers.
|
||
code := tpl.Server
|
||
imports := []string{
|
||
"\"" + moduleName + "/internal/logic/" + toSnakeCase(service.GoName) + "\"",
|
||
"pb \"" + moduleName + "/pb\"",
|
||
}
|
||
|
||
code = strings.ReplaceAll(code, "{import}", strings.Join(imports, "\n"))
|
||
code = strings.ReplaceAll(code, "{service}", service.GoName)
|
||
|
||
var codeMethods []string
|
||
for _, method := range service.Methods {
|
||
commit := strings.TrimSpace(method.Comments.Leading.String())
|
||
methodCode := tpl.Method
|
||
methodCode = strings.ReplaceAll(methodCode, "{service}", service.GoName)
|
||
methodCode = strings.ReplaceAll(methodCode, "{serviceLower}", strings.ToLower(service.GoName))
|
||
methodCode = strings.ReplaceAll(methodCode, "{func}", method.GoName)
|
||
methodCode = strings.ReplaceAll(methodCode, "{comment}", commit)
|
||
methodCode = strings.ReplaceAll(methodCode, "{input}", method.Input.GoIdent.GoName)
|
||
methodCode = strings.ReplaceAll(methodCode, "{output}", method.Output.GoIdent.GoName)
|
||
codeMethods = append(codeMethods, methodCode)
|
||
}
|
||
code = strings.ReplaceAll(code, "{method}", strings.Join(codeMethods, "\n"))
|
||
|
||
// 格式化代码
|
||
formattedCode, err := format.Source([]byte(code))
|
||
if err != nil {
|
||
return fmt.Errorf("failed to format generated code: %w", err)
|
||
}
|
||
|
||
StringToFile(filename, string(formattedCode))
|
||
|
||
return nil
|
||
}
|
||
|
||
func generateLogicFile(gen *protogen.Plugin, file *protogen.File, service *protogen.Service) error {
|
||
logicPath := "./internal/logic/" + toSnakeCase(service.GoName)
|
||
if !utils.PathExists(logicPath) {
|
||
os.MkdirAll(logicPath, os.ModePerm)
|
||
}
|
||
moduleName := getModuleName()
|
||
for _, method := range service.Methods {
|
||
filename := fmt.Sprintf("%s/%s.go", logicPath, toSnakeCase(method.GoName))
|
||
if utils.PathExists(filename) {
|
||
continue
|
||
}
|
||
code := tpl.LogicFile
|
||
|
||
code = strings.ReplaceAll(code, "{methodName}", strings.ToLower(service.GoName))
|
||
imports := []string{
|
||
"pb \"" + moduleName + "/pb\"",
|
||
}
|
||
|
||
if strings.ToLower(method.Input.GoIdent.GoName) == "identrequest" || strings.ToLower(method.Input.GoIdent.GoName) == "fetchrequest" {
|
||
if strings.ToLower(method.Input.GoIdent.GoName) == "identrequest" {
|
||
imports = append(imports, "\"git.apinb.com/bsm-sdk/core/errcode\"")
|
||
code = strings.ReplaceAll(code, "{valid}", tpl.ValidCode)
|
||
}
|
||
if strings.ToLower(method.Input.GoIdent.GoName) == "fetchrequest" {
|
||
code = strings.ReplaceAll(code, "{valid}", tpl.FetchValidCode)
|
||
}
|
||
|
||
} else {
|
||
code = strings.ReplaceAll(code, "{valid}", "// TODO: valid code")
|
||
}
|
||
|
||
if strings.ToLower(method.Output.GoIdent.GoName) == "statusreply" {
|
||
imports = append(imports, "\"time\"")
|
||
imports = append(imports, "\"git.apinb.com/bsm-sdk/core/vars\"")
|
||
code = strings.ReplaceAll(code, "{return}", tpl.StatusReplyCode)
|
||
} else {
|
||
code = strings.ReplaceAll(code, "{return}", "return ")
|
||
}
|
||
|
||
code = strings.ReplaceAll(code, "{import}", strings.Join(imports, "\n"))
|
||
commit := strings.TrimSpace(method.Comments.Leading.String())
|
||
code = strings.ReplaceAll(code, "{func}", method.GoName)
|
||
code = strings.ReplaceAll(code, "{comment}", commit)
|
||
code = strings.ReplaceAll(code, "{input}", method.Input.GoIdent.GoName)
|
||
code = strings.ReplaceAll(code, "{output}", method.Output.GoIdent.GoName)
|
||
|
||
formattedCode, err := format.Source([]byte(code))
|
||
if err != nil {
|
||
return fmt.Errorf("failed to format generated code: %w", err)
|
||
}
|
||
|
||
StringToFile(filename, string(formattedCode))
|
||
}
|
||
|
||
return nil
|
||
}
|
||
|
||
func toSnakeCase(str string) string {
|
||
// Use a regular expression to find uppercase letters and insert an underscore before them
|
||
re := regexp.MustCompile("([a-z0-9])([A-Z])")
|
||
snake := re.ReplaceAllString(str, "${1}_${2}")
|
||
// Convert the entire string to lowercase
|
||
return strings.ToLower(snake)
|
||
}
|
||
|
||
func methodSignature(g *protogen.GeneratedFile, method *protogen.Method) string {
|
||
return fmt.Sprintf("%s(ctx context.Context, req pb%s) (*%s, error)",
|
||
method.GoName,
|
||
method.Input.GoIdent,
|
||
method.Output.GoIdent)
|
||
}
|
||
|
||
func fullMethodName(file *protogen.File, service *protogen.Service, method *protogen.Method) string {
|
||
return fmt.Sprintf("/%s.%s/%s",
|
||
file.Proto.GetPackage(),
|
||
service.GoName,
|
||
method.GoName)
|
||
}
|
||
|
||
func getModuleName() (modulePath string) {
|
||
// 获取当前工作目录
|
||
cwd, err := os.Getwd()
|
||
if err != nil {
|
||
fmt.Errorf("failed to get current working directory: %w", err)
|
||
return
|
||
}
|
||
|
||
// 读取 go.mod 文件
|
||
modFilePath := filepath.Join(cwd, "go.mod")
|
||
modFileBytes, err := os.ReadFile(modFilePath)
|
||
if err != nil {
|
||
fmt.Errorf("failed to read go.mod file: %w", err)
|
||
return
|
||
}
|
||
|
||
// 解析 go.mod 文件
|
||
modFile, err := modfile.Parse(modFilePath, modFileBytes, nil)
|
||
if err != nil {
|
||
fmt.Errorf("failed to parse go.mod file: %w", err)
|
||
return
|
||
}
|
||
|
||
// 获取模块路径
|
||
return modFile.Module.Mod.Path
|
||
}
|
||
|
||
// 将字符串写入文件
|
||
func StringToFile(path, content string) error {
|
||
startF, err := os.Create(path)
|
||
if err != nil {
|
||
return errors.New("os.Create create file " + path + " error:" + err.Error())
|
||
}
|
||
defer startF.Close()
|
||
_, err = io.WriteString(startF, content)
|
||
if err != nil {
|
||
return errors.New("io.WriteString to " + path + " error:" + err.Error())
|
||
}
|
||
return nil
|
||
}
|
||
|
||
// parseOptions 解析注释中的选项字符串并返回一个 map
|
||
func parseOptions(comment string) map[string]string {
|
||
// 去掉注释符号和分号
|
||
comment = strings.Trim(comment, " //;")
|
||
// 按逗号分割选项
|
||
options := strings.Split(comment, ",")
|
||
result := make(map[string]string)
|
||
|
||
for _, opt := range options {
|
||
// 按等号分割键和值
|
||
parts := strings.SplitN(opt, "=", 2)
|
||
if len(parts) == 2 {
|
||
key := strings.TrimSpace(parts[0])
|
||
value := strings.TrimSpace(parts[1])
|
||
result[key] = value
|
||
}
|
||
}
|
||
|
||
return result
|
||
}
|
||
|
||
// CamelToSnake 将驼峰命名转换为下划线命名
|
||
func CamelToSnake(s string) string {
|
||
var buf bytes.Buffer
|
||
for i, r := range s {
|
||
if unicode.IsUpper(r) {
|
||
// 如果不是第一个字符,添加下划线
|
||
if i > 0 {
|
||
buf.WriteRune('_')
|
||
}
|
||
buf.WriteRune(unicode.ToLower(r))
|
||
} else {
|
||
buf.WriteRune(r)
|
||
}
|
||
}
|
||
return buf.String()
|
||
}
|