From 4bf991b3e720cd9155fd0983c64fa3bdc699b747 Mon Sep 17 00:00:00 2001 From: chiyi Date: Sun, 16 Aug 2026 23:51:58 +0800 Subject: [PATCH] fun: HTTP RPC framework with DI, guards, streams and codegen - Service binding via reflection with 4 method signatures - Dependency injection (Wired/auto tags) and guards - Streamable HTTP (NDJSON) streaming responses - Go/TypeScript client code generation (self-contained, with interceptors) - Enum support, structured Result, logging - 16 regression tests covering all fixed bugs --- .gitignore | 23 +++ box.go | 159 ++++++++++++++++++++ bugfix_test.go | 222 ++++++++++++++++++++++++++++ check.go | 165 +++++++++++++++++++++ ctx.go | 86 +++++++++++ enum.go | 26 ++++ fun.go | 156 ++++++++++++++++++++ fun_test.go | 159 ++++++++++++++++++++ gen.go | 176 +++++++++++++++++++++++ gen_go.go | 222 ++++++++++++++++++++++++++++ gen_ts.go | 250 ++++++++++++++++++++++++++++++++ go.mod | 11 ++ go.sum | 10 ++ guard.go | 5 + handle.go | 226 +++++++++++++++++++++++++++++ logger.go | 384 +++++++++++++++++++++++++++++++++++++++++++++++++ request.go | 14 ++ result.go | 60 ++++++++ stream.go | 80 +++++++++++ tag.go | 34 +++++ template_go.go | 278 +++++++++++++++++++++++++++++++++++ template_ts.go | 186 ++++++++++++++++++++++++ text.go | 25 ++++ utils.go | 60 ++++++++ 24 files changed, 3017 insertions(+) create mode 100644 .gitignore create mode 100644 box.go create mode 100644 bugfix_test.go create mode 100644 check.go create mode 100644 ctx.go create mode 100644 enum.go create mode 100644 fun.go create mode 100644 fun_test.go create mode 100644 gen.go create mode 100644 gen_go.go create mode 100644 gen_ts.go create mode 100644 go.mod create mode 100644 go.sum create mode 100644 guard.go create mode 100644 handle.go create mode 100644 logger.go create mode 100644 request.go create mode 100644 result.go create mode 100644 stream.go create mode 100644 tag.go create mode 100644 template_go.go create mode 100644 template_ts.go create mode 100644 text.go create mode 100644 utils.go diff --git a/.gitignore b/.gitignore new file mode 100644 index 0000000..d20a2dc --- /dev/null +++ b/.gitignore @@ -0,0 +1,23 @@ +# IDE +.idea/ +*.iml + +# Binaries +*.exe +*.dll +*.so +*.dylib + +# Generated code +gen/ + +# Logs +*.log + +# OS files +.DS_Store +Thumbs.db + +# Dependencies +vendor/ +node_modules/ diff --git a/box.go b/box.go new file mode 100644 index 0000000..0dbcb08 --- /dev/null +++ b/box.go @@ -0,0 +1,159 @@ +package fun + +import ( + "reflect" +) + +// Wired 创建并注册一个依赖实例;auto 标签字段递归注入依赖;存在 New() 则调用 +func Wired[T any]() *T { + t := reflect.TypeFor[T]() + data := new(T) + if t.Kind() != reflect.Struct { + panic("Fun: " + t.Name() + " It must be a structure") + } + if isPrivate(t.Name()) { + panic("Fun:" + t.Name() + " cannot be Private") + } + if newMethod, found := t.MethodByName("New"); found { + if newMethod.Type.NumIn() != 1 || newMethod.Type.NumOut() != 0 { + panic("Fun:" + t.Name() + " New method must have no parameters and no return values") + } + } + f := GetFun() + if box, isWired := f.boxes.Load(reflect.TypeFor[*T]()); isWired { + return box.(reflect.Value).Interface().(*T) + } + v := reflect.ValueOf(data) + f.boxes.Store(reflect.TypeFor[*T](), v) + boxList := map[reflect.Type]bool{} + for i := 0; i < t.NumField(); i++ { + c := t.Field(i) + fieldTag := newTag(c.Tag) + if _, isAuto := fieldTag.getTag("auto"); isAuto { + if dependency, loaded := f.boxes.Load(c.Type); loaded { + v.Elem().Field(i).Set(dependency.(reflect.Value)) + } else { + checkBox(c, boxList) + f.autowired(v.Elem().Field(i)) + } + } + } + newMethod := v.MethodByName("New") + if newMethod.IsValid() { + newMethod.Call(nil) + } + return data +} + +// autowired 递归创建依赖实例并注入 auto 标签字段 +func (f *Fun) autowired(fieldValue reflect.Value) { + instance := reflect.New(fieldValue.Type().Elem()) + f.boxes.Store(fieldValue.Type(), instance) + fieldValue.Set(instance) + structValue := instance.Elem() + for i := 0; i < structValue.NumField(); i++ { + structField := structValue.Type().Field(i) + fieldTag := newTag(structField.Tag) + if _, isAuto := fieldTag.getTag("auto"); isAuto { + if dependency, loaded := f.boxes.Load(structField.Type); loaded { + structValue.Field(i).Set(dependency.(reflect.Value)) + } else { + f.autowired(structValue.Field(i)) + } + } + } + newMethod := instance.MethodByName("New") + if newMethod.IsValid() { + newMethod.Call(nil) + } +} + +// checkBox 校验 auto 注入字段:必须是指针+struct、非匿名、非私有;New() 必须无参无返回值 +func checkBox(s reflect.StructField, boxList map[reflect.Type]bool) { + if _, ok := boxList[s.Type]; ok { + return + } + boxList[s.Type] = true + if s.Anonymous { + panic("Fun:" + s.Name + " cannot be Anonymous") + } + if s.Type.Kind() != reflect.Ptr || s.Type.Elem().Kind() != reflect.Struct { + panic("Fun:" + s.Name + " Must be a pointer and a struct") + } + if isPrivate(s.Name) { + panic("Fun:" + s.Name + " cannot be Private") + } + if newMethod, found := s.Type.MethodByName("New"); found { + if newMethod.Type.NumIn() != 1 || newMethod.Type.NumOut() != 0 { + panic("Fun:" + s.Name + " New method must have no parameters and no return values") + } + } + for i := 0; i < s.Type.Elem().NumField(); i++ { + f := s.Type.Elem().Field(i) + fieldTag := newTag(f.Tag) + if _, isAuto := fieldTag.getTag("auto"); isAuto { + checkBox(f, boxList) + } + } +} + +// boxWired 注册期预初始化服务结构体字段中的 Box 依赖 +func boxWired(service any, f *Fun) { + serviceInstance := reflect.New(reflect.TypeOf(service).Elem()).Elem() + for i := 0; i < serviceInstance.NumField(); i++ { + field := serviceInstance.Field(i) + if field.Type() == ctxType { + continue + } + if field.Type().Kind() == reflect.Ptr && field.Type().Elem().Kind() == reflect.Struct { + if _, isWired := f.boxes.Load(field.Type()); !isWired { + f.autowired(field) + } + } + } +} + +// serviceWired 每请求把 Ctx 与 Box 依赖注入到新创建的服务实例 +func (f *Fun) serviceWired(serviceInstance reflect.Value, ctx *Ctx) { + for i := 0; i < serviceInstance.NumField(); i++ { + field := serviceInstance.Field(i) + if !field.CanSet() { + continue + } + if field.Type() == ctxType { + field.Set(reflect.ValueOf(*ctx)) + } else if dependency, ok := f.boxes.Load(field.Type()); ok { + field.Set(dependency.(reflect.Value)) + } + } +} + +// checkGuard 校验 Guard 类型:必须是指向结构体的指针 +func checkGuard(guard Guard) { + t := reflect.TypeOf(guard) + if t.Kind() != reflect.Ptr || t.Elem().Kind() != reflect.Struct { + panic("Fun: guard must be a pointer to a struct") + } + if isPrivate(t.Elem().Name()) { + panic("Fun:" + t.Elem().Name() + " cannot be Private") + } +} + +// serviceGuardWired 创建 Guard 实例并注入 Box 依赖,返回 guard 引用 +func serviceGuardWired(guard Guard, f *Fun) *any { + t := reflect.TypeOf(guard).Elem() + guardInstance := reflect.New(t).Elem() + for i := 0; i < guardInstance.NumField(); i++ { + field := guardInstance.Field(i) + if !field.CanSet() { + continue + } + if dependency, ok := f.boxes.Load(field.Type()); ok { + field.Set(dependency.(reflect.Value)) + } else { + f.autowired(field) + } + } + g := guardInstance.Addr().Interface() + return &g +} diff --git a/bugfix_test.go b/bugfix_test.go new file mode 100644 index 0000000..dcb55e7 --- /dev/null +++ b/bugfix_test.go @@ -0,0 +1,222 @@ +package fun + +import ( + "bufio" + "bytes" + "encoding/json" + "errors" + "net/http" + "os" + "path/filepath" + "strings" + "testing" + "time" +) + +type BugStatus uint8 + +func (BugStatus) Names() []string { return []string{"A", "B"} } + +type BugDto struct { + Status *BugStatus +} + +type BugSvc struct{} + +func (s *BugSvc) Ping() error { return nil } + +func (s *BugSvc) Save(dto BugDto) (string, error) { return "ok", nil } + +func (s *BugSvc) Ticker() (string, *Stream, error) { + st := Stream{} + go func() { + st.Send("tick") + st.Close() + }() + return "first", &st, nil +} + +func bugInvoke(t *testing.T, method string, data map[string]any) (*Result[any], error) { + t.Helper() + f := New() + f.BindService(&BugSvc{}) + if data == nil { + data = map[string]any{} + } + c := &Ctx{Ip: "1", MethodName: method, ServiceName: "BugSvc", Data: &data} + var streamCh chan any + var streamDone chan struct{} + return f.invoke(c, &streamCh, &streamDone) +} + +// bug1: () error 签名的方法 invoke 应返回空数据结果而不是越界 panic +func TestBugErrorOnlyInvoke(t *testing.T) { + res, err := bugInvoke(t, "Ping", nil) + if err != nil { + t.Fatalf("err: %v", err) + } + if res.Data != nil { + t.Fatalf("expect nil data, got %v", *res.Data) + } +} + +// bug2: 指针枚举字段传 null 应放行;越界值仍要报错 +func TestBugNullableEnum(t *testing.T) { + if _, err := bugInvoke(t, "Save", map[string]any{"status": nil}); err != nil { + t.Fatalf("nullable enum should pass: %v", err) + } + if _, err := bugInvoke(t, "Save", map[string]any{"status": 1}); err != nil { + t.Fatalf("valid enum should pass: %v", err) + } + if _, err := bugInvoke(t, "Save", map[string]any{"status": 5}); err == nil { + t.Fatal("out-of-range enum should fail") + } +} + +// bug3: 含 () error 方法的代码生成不应 panic,且类型应生成为 Void/void +func TestBugGenErrorOnly(t *testing.T) { + GetFun().BindService(&BugSvc{}) + SetOutput(t.TempDir()) + GenCode(GenGo{}, GenTs{}) + goSrc, err := os.ReadFile(filepath.Join(getDirectory(), "go", "bug_svc.go")) + if err != nil { + t.Fatalf("go file missing: %v", err) + } + if !strings.Contains(string(goSrc), "Result[Void]") { + t.Fatalf("go: expect Result[Void], got:\n%s", goSrc) + } + tsSrc, err := os.ReadFile(filepath.Join(getDirectory(), "ts", "bugSvc.ts")) + if err != nil { + t.Fatalf("ts file missing: %v", err) + } + if !strings.Contains(string(tsSrc), "result") { + t.Fatalf("ts: expect result, got:\n%s", tsSrc) + } +} + +// bug4+5: 响应键应为小写;(T, stream, error) 的 T 应作为流的第一条消息下发 +func TestBugJsonKeysAndStreamFirst(t *testing.T) { + f := New() + f.BindService(&BugSvc{}) + go f.Start(39003) + time.Sleep(300 * time.Millisecond) + + resp, err := http.Post("http://127.0.0.1:39003/cell", "application/json", + strings.NewReader(`{"serviceName":"BugSvc","methodName":"Ping"}`)) + if err != nil { + t.Fatal(err) + } + var buf bytes.Buffer + _, _ = buf.ReadFrom(resp.Body) + _ = resp.Body.Close() + body := buf.String() + if !strings.Contains(body, `"status"`) || strings.Contains(body, `"Status"`) { + t.Fatalf("keys not lowercase: %s", body) + } + + resp2, err := http.Post("http://127.0.0.1:39003/cell", "application/json", + bytes.NewReader([]byte(`{"serviceName":"BugSvc","methodName":"Ticker"}`))) + if err != nil { + t.Fatal(err) + } + defer resp2.Body.Close() + var got []string + scanner := bufio.NewScanner(resp2.Body) + for scanner.Scan() { + line := strings.TrimSpace(scanner.Text()) + if line == "" { + continue + } + var msg string + if err := json.Unmarshal([]byte(line), &msg); err != nil { + t.Fatalf("bad ndjson line %q: %v", line, err) + } + got = append(got, msg) + } + if len(got) != 2 || got[0] != "first" || got[1] != "tick" { + t.Fatalf("stream: %v", got) + } +} + +// bug7: 非指针切片字段传 null 应放行(JSON null -> nil slice) +type NullSlicDto struct { + Tags []string +} + +type NullSlicSvc struct{} + +func (s *NullSlicSvc) Save(dto NullSlicDto) (string, error) { return "ok", nil } + +func TestBugSliceNull(t *testing.T) { + f := New() + f.BindService(&NullSlicSvc{}) + data := map[string]any{"tags": nil} + c := &Ctx{Ip: "1", MethodName: "Save", ServiceName: "NullSlicSvc", Data: &data} + var streamCh chan any + var streamDone chan struct{} + if _, err := f.invoke(c, &streamCh, &streamDone); err != nil { + t.Fatalf("slice null should pass: %v", err) + } +} + +// bug8: 业务方法返回 error 但已启动 goroutine 调 Send,框架应注入取消流让 goroutine 退出而非挂死 +var leakDone chan struct{} + +type LeakSvc struct{} + +func (s *LeakSvc) Fail() (*Stream, error) { + st := Stream{} + leakDone = make(chan struct{}) + go func() { + st.Send("never") + close(leakDone) + }() + return &st, errors.New("boom") +} + +func TestBugStreamLeak(t *testing.T) { + f := New() + f.BindService(&LeakSvc{}) + c := &Ctx{Ip: "1", MethodName: "Fail", ServiceName: "LeakSvc"} + var streamCh chan any + var streamDone chan struct{} + _, err := f.invoke(c, &streamCh, &streamDone) + if err == nil { + t.Fatal("expected error") + } + // Send goroutine 必须解除阻塞(有超时保护,防挂死) + select { + case <-leakDone: + case <-time.After(2 * time.Second): + t.Fatal("Send goroutine blocked forever: stream leak") + } +} + +// bug6: 与 fasthttp.RequestCtx 方法重名的用户方法不应被静默丢弃 +type CollideSvc struct { + Ctx +} + +func (s *CollideSvc) Cookie() (string, error) { return "cookie", nil } + +func TestBugMethodNameCollision(t *testing.T) { + f := New() + f.BindService(&CollideSvc{}) + if _, ok := f.methods["CollideSvc.Cookie"]; !ok { + t.Fatal("Cookie method dropped due to name collision with fasthttp.RequestCtx") + } + res, err := func() (*Result[any], error) { + c := &Ctx{Ip: "1", MethodName: "Cookie", ServiceName: "CollideSvc"} + data := map[string]any{} + c.Data = &data + var streamCh chan any + var streamDone chan struct{} + return f.invoke(c, &streamCh, &streamDone) + }() + if err != nil { + t.Fatalf("invoke Cookie err: %v", err) + } + if (*res.Data).(string) != "cookie" { + t.Fatalf("unexpected: %v", *res.Data) + } +} diff --git a/check.go b/check.go new file mode 100644 index 0000000..121ed53 --- /dev/null +++ b/check.go @@ -0,0 +1,165 @@ +package fun + +import ( + "errors" + "fmt" + "reflect" + "strings" + "unicode" +) + +func isPrivate(value string) bool { + return !unicode.IsUpper([]rune(value)[0]) +} + +// checkType 注册期递归校验类型是否受支持: +// int/uint/string/bool/struct/slice/enum;不支持匿名结构体、私有类型、空结构体; +// 枚举要求 Names/DisplayNames 长度一致 +func checkType(t reflect.Type) { + if t.Kind() == reflect.Ptr { + t = t.Elem() + } + if strings.Contains(t.String(), "{}") { + panic(fmt.Sprintf("fun: %s generic types containing 'any' or interface{} are not supported", t.Name())) + } + switch t.Kind() { + case reflect.Int8, reflect.Int16, reflect.Int32, reflect.Int64, + reflect.Uint8, reflect.Uint16, reflect.Uint32, reflect.Uint64, + reflect.String, reflect.Bool: + if t.Kind() == reflect.Uint8 && (t.Implements(displayEnumType) || t.Implements(enumType)) && isPrivate(t.Name()) { + panic("fun:" + t.Name() + " cannot be Private") + } + if t.Kind() == reflect.Uint8 && t.Implements(displayEnumType) { + enumValue := reflect.New(t).Elem().Interface().(displayEnum) + if len(enumValue.DisplayNames()) != len(enumValue.Names()) { + panic("fun: " + t.Name() + " enum names and display names must be the same length") + } + } + case reflect.Struct: + if t.NumField() == 0 { + panic("fun: " + t.Name() + " must have at least one field") + } + for i := 0; i < t.NumField(); i++ { + f := t.Field(i) + if isPrivate(f.Name) { + panic("fun:" + f.Name + " cannot be Private") + } + checkType(f.Type) + } + case reflect.Slice: + checkType(t.Elem()) + default: + panic("fun:Unsupported types " + t.Name()) + } +} + +// checkDto 运行时校验请求数据: +// 非指针字段必须出现在请求中且非 nil;嵌套 struct/slice 递归;枚举值必须在范围内 +func checkDto(dtoType reflect.Type, dtoMap any, methodName string) error { + t := dtoType + if t.Kind() == reflect.Ptr { + t = t.Elem() + } + switch t.Kind() { + case reflect.Struct: + obj, ok := dtoMap.(map[string]any) + if !ok { + return callError(fmt.Errorf("fun: method %s DTO %s must be an object", methodName, t.Name())) + } + for i := 0; i < t.NumField(); i++ { + f := t.Field(i) + value, ok := obj[firstLetterToLower(f.Name)] + if !ok { + // 兼容按原始字段名(首字母大写)传参的客户端 + value, ok = obj[f.Name] + } + // 非指针且非切片字段必须存在且非 null;切片字段允许 null(反序列化为 nil slice) + if f.Type.Kind() != reflect.Ptr && f.Type.Kind() != reflect.Slice && (!ok || value == nil) { + return callError(fmt.Errorf("fun: %s Dto must be a pointer or have a corresponding field in the map", f.Name)) + } + ft := f.Type + if ft.Kind() == reflect.Ptr { + ft = ft.Elem() + } + if (ft.Kind() == reflect.Struct || ft.Kind() == reflect.Slice) && value != nil { + if err := checkDto(ft, value, methodName); err != nil { + return err + } + } + if ft.Kind() == reflect.Uint8 && value != nil && (ft.Implements(displayEnumType) || ft.Implements(enumType)) { + if err := checkEnumValue(ft, value, f.Name); err != nil { + return err + } + } + } + case reflect.Slice: + list, ok := dtoMap.([]any) + if !ok { + return callError(fmt.Errorf("fun: Dto must be an array")) + } + for _, value := range list { + et0 := t.Elem() + et := et0 + if et.Kind() == reflect.Ptr { + et = et.Elem() + } + // 指针元素允许 null,值元素必须非空 + if et0.Kind() != reflect.Ptr && value == nil { + return callError(fmt.Errorf("fun:%s Dto must be a pointer or have a corresponding field in the map", et0.Name())) + } + if (et.Kind() == reflect.Struct || et.Kind() == reflect.Slice) && value != nil { + if err := checkDto(et, value, methodName); err != nil { + return err + } + } + if et.Kind() == reflect.Uint8 && value != nil && (et.Implements(displayEnumType) || et.Implements(enumType)) { + if err := checkEnumValue(et, value, et.Name()); err != nil { + return err + } + } + } + } + return nil +} + +// checkEnumValue 运行时校验枚举值是否在范围内 +func checkEnumValue(t reflect.Type, value any, name string) error { + var max uint8 + enumValue := reflect.New(t).Elem() + if t.Implements(displayEnumType) { + max = uint8(len(enumValue.Interface().(displayEnum).Names())) + } else { + max = uint8(len(enumValue.Interface().(enum).Names())) + } + var num uint8 + switch v := value.(type) { + case float64: + num = uint8(v) + case float32: + num = uint8(v) + case uint8: + num = v + case uint16: + num = uint8(v) + case uint32: + num = uint8(v) + case uint64: + num = uint8(v) + case int: + num = uint8(v) + case int8: + num = uint8(v) + case int16: + num = uint8(v) + case int32: + num = uint8(v) + case int64: + num = uint8(v) + default: + return callError(errors.New("Fun:" + name + " Dto enum value type is not supported")) + } + if num >= max { + return callError(errors.New("Fun:" + name + " Dto value out of range")) + } + return nil +} diff --git a/ctx.go b/ctx.go new file mode 100644 index 0000000..987bc2d --- /dev/null +++ b/ctx.go @@ -0,0 +1,86 @@ +package fun + +import ( + "encoding/json" + "net" + "reflect" + + "github.com/valyala/fasthttp" +) + +// Ctx 请求上下文。 +// 命名持有 *fasthttp.RequestCtx(非嵌入),避免其方法集提升到服务上, +// 辅助方法全部小写,保证服务方法集只含业务方法。 +type Ctx struct { + Ip string + State map[string]string + MethodName string + ServiceName string + Data *map[string]any + RequestCtx *fasthttp.RequestCtx +} + +var ctxType = reflect.TypeFor[Ctx]() + +func (c *Ctx) path() string { return string(c.RequestCtx.Path()) } +func (c *Ctx) isPost() bool { return c.RequestCtx.IsPost() } +func (c *Ctx) postBody() []byte { return c.RequestCtx.PostBody() } +func (c *Ctx) remoteIP() net.IP { return c.RequestCtx.RemoteIP() } +func (c *Ctx) setStatusCode(code int) { c.RequestCtx.SetStatusCode(code) } +func (c *Ctx) write(p []byte) (int, error) { return c.RequestCtx.Write(p) } + +// send 写回响应,嵌套对象键统一转小写(兼容 TS/大小写敏感客户端) +func (c *Ctx) send(result Result[any]) { + data, err := json.Marshal(result) + if err != nil { + c.sendError(err) + return + } + raw, err := lowerKeysFromJSON(data) + if err != nil { + _, _ = c.write(data) + return + } + out, err := json.Marshal(raw) + if err != nil { + _, _ = c.write(data) + return + } + _, _ = c.write(out) +} + +// lowerKeysFromJSON 解析 JSON 后递归把所有对象键转为首字母小写 +func lowerKeysFromJSON(data []byte) (any, error) { + var raw any + if err := json.Unmarshal(data, &raw); err != nil { + return nil, err + } + return lowerKeys(raw), nil +} + +func lowerKeys(obj any) any { + switch v := obj.(type) { + case map[string]any: + m := make(map[string]any, len(v)) + for k, val := range v { + m[firstLetterToLower(k)] = lowerKeys(val) + } + return m + case []any: + for i := range v { + v[i] = lowerKeys(v[i]) + } + return v + } + return obj +} + +// sendError 写回错误响应 +// 业务 Error() 构造的 Result[any] 原样透传保留 Code/Msg/Status,普通 error 包成错误响应 +func (c *Ctx) sendError(err error) { + if result, ok := err.(Result[any]); ok { + c.send(result) + return + } + c.send(callError(err)) +} diff --git a/enum.go b/enum.go new file mode 100644 index 0000000..a48578b --- /dev/null +++ b/enum.go @@ -0,0 +1,26 @@ +package fun + +import ( + "errors" + "reflect" +) + +type enum interface { + Names() []string +} + +type displayEnum interface { + DisplayNames() []string + Names() []string +} + +var ( + enumType = reflect.TypeFor[enum]() + displayEnumType = reflect.TypeFor[displayEnum]() +) + +var ( + errMethodNotFound = errors.New("method not found") + errEmptyFields = errors.New("serviceName and methodName cannot be empty") + errDTORequired = errors.New("method requires a DTO but none provided") +) diff --git a/fun.go b/fun.go new file mode 100644 index 0000000..a780e18 --- /dev/null +++ b/fun.go @@ -0,0 +1,156 @@ +package fun + +import ( + "fmt" + "reflect" + "sync" + + "github.com/valyala/fasthttp" +) + +type Fun struct { + methods map[string]methodInfo + boxes *sync.Map // 依赖容器:reflect.Type → reflect.Value + guards []*any // 全局 Guard + serviceGuards map[string][]*any // 服务级 Guard,按服务名 +} + +var ( + errorType = reflect.TypeFor[error]() + streamType = reflect.TypeFor[*Stream]() +) + +var fun *Fun + +// methodInfo 已注册方法的元信息 +type methodInfo struct { + serviceType reflect.Type // 服务值类型(非指针),每请求新建实例 + methodIndex int // 方法在实例上的反射索引 + dtoType reflect.Type // DTO 参数类型,无参数时为 nil + isStream bool // 返回签名带 *Stream,走 RequestStreamType +} + +func New() *Fun { + f := &Fun{ + methods: map[string]methodInfo{}, + boxes: &sync.Map{}, + serviceGuards: map[string][]*any{}, + } + if fun == nil { + fun = f + } + return f +} + +// GetFun 返回默认 Fun 实例,未初始化时自动创建 +func GetFun() *Fun { + if fun == nil { + fun = New() + } + return fun +} + +// BindService 注册服务,要求传入指向结构体的指针 +// 方法签名约束: +// - 参数:最多一个,且必须是 struct(作为 DTO) +// - 返回值:只支持四种签名——(error)、(T, error)、(stream, error)、(T, stream, error) +// +// guardList 为该服务绑定的 Guard,方法调用前按注册顺序执行 +func (f *Fun) BindService(service any, guardList ...Guard) { + t := reflect.TypeOf(service) + // 必须是指针指向的结构体,匿名类型无法注册 + if t.Kind() != reflect.Ptr || t.Elem().Kind() != reflect.Struct { + panic("fun: BindService requires a pointer to a struct") + } + name := t.Elem().Name() + if name == "" { + panic("fun: BindService requires a named type") + } + + boxWired(service, f) + + serviceGuards := make([]*any, 0, len(guardList)) + for _, guard := range guardList { + checkGuard(guard) + serviceGuards = append(serviceGuards, serviceGuardWired(guard, f)) + } + f.serviceGuards[name] = serviceGuards + + for m := range t.Methods() { + m := m + // Ctx 命名持有 *fasthttp.RequestCtx(非嵌入),服务方法集只含业务方法,无需过滤提升方法 + mt := m.Type + + // 参数:接收者 + 最多一个 DTO(NumIn() 含接收者),DTO 必须是 struct + if mt.NumIn() > 2 { + panic(fmt.Sprintf("fun: method %s has more than one parameter", m.Name)) + } + var dtoType reflect.Type + if mt.NumIn() == 2 { + dtoType = mt.In(1) + if dtoType.Kind() != reflect.Struct { + panic(fmt.Sprintf("fun: method %s parameter must be a struct", m.Name)) + } + checkType(dtoType) + } + + // 返回值只支持四种签名:error / (T, error) / (stream, error) / (T, stream, error) + isStream := false + switch mt.NumOut() { + case 1: + // 情况 1:func(...) error + if mt.Out(0) != errorType { + panic(fmt.Sprintf("fun: method %s must return (error), (T, error), (stream, error) or (T, stream, error)", m.Name)) + } + case 2: + // 情况 2:func(...) (T, error) 或 func(...) (*Stream, error) + if mt.Out(1) != errorType { + panic(fmt.Sprintf("fun: method %s last return value must be error", m.Name)) + } + isStream = mt.Out(0) == streamType + case 3: + // 情况 3:func(...) (T, *Stream, error) + if mt.Out(2) != errorType { + panic(fmt.Sprintf("fun: method %s last return value must be error", m.Name)) + } + if mt.Out(1) != streamType { + panic(fmt.Sprintf("fun: method %s second return value must be *Stream", m.Name)) + } + isStream = true + default: + panic(fmt.Sprintf("fun: method %s must return (error), (T, error), (stream, error) or (T, stream, error)", m.Name)) + } + + // 注册到 "ServiceName.MethodName" + f.methods[name+"."+m.Name] = methodInfo{ + serviceType: t.Elem(), + methodIndex: m.Index, + dtoType: dtoType, + isStream: isStream, + } + } +} + +// BindGuard 注册全局 Guard,对所有服务生效 +func (f *Fun) BindGuard(guard Guard) { + checkGuard(guard) + f.guards = append(f.guards, serviceGuardWired(guard, f)) +} + +// callGuard 按 全局 → 服务级 顺序执行 Guard +func (f *Fun) callGuard(c *Ctx, serviceName string) { + for _, g := range f.guards { + (*g).(Guard).Guard(*c) + } + for _, g := range f.serviceGuards[serviceName] { + (*g).(Guard).Guard(*c) + } +} + +func (f *Fun) Start(port uint16) { + addr := fmt.Sprintf(":%d", port) + err := fasthttp.ListenAndServe(addr, f.handle) + if err != nil { + panic(err.Error()) + } +} diff --git a/fun_test.go b/fun_test.go new file mode 100644 index 0000000..455d720 --- /dev/null +++ b/fun_test.go @@ -0,0 +1,159 @@ +package fun + +import ( + "bufio" + "bytes" + "encoding/json" + "errors" + "fmt" + "net/http" + "os" + "path/filepath" + "testing" + "time" +) + +type TestRepo struct{} + +type TestDto struct { + Name string + Age uint8 +} + +type TestSvc struct { + Ctx + Repo *TestRepo +} + +var guardHit = false + +type TestGuard struct{} + +func (g *TestGuard) Guard(ctx Ctx) { + guardHit = true +} + +func (s *TestSvc) Hello(dto TestDto) (string, error) { + if s.Ip == "" { + return "", errors.New("no ip") + } + if s.Repo == nil { + return "", errors.New("no repo") + } + return "hi " + dto.Name, nil +} + +func (s *TestSvc) Count(dto TestDto) (*Stream, error) { + st := Stream{} + go func() { + for i := 0; i < 3; i++ { + st.Send(fmt.Sprintf("n%d", i)) + time.Sleep(5 * time.Millisecond) + } + st.Close() + }() + return &st, nil +} + +func TestCtxBoxInject(t *testing.T) { + f := New() + guardHit = false + f.BindService(&TestSvc{}, &TestGuard{}) + c := &Ctx{Ip: "1.2.3.4", MethodName: "Hello", ServiceName: "TestSvc"} + data := map[string]any{"name": "tom", "age": 1} + c.Data = &data + + var streamCh chan any + var streamDone chan struct{} + res, err := f.invoke(c, &streamCh, &streamDone) + if err != nil { + t.Fatalf("invoke err: %v", err) + } + if (*res.Data).(string) != "hi tom" { + t.Fatalf("unexpected data: %v", *res.Data) + } + if !guardHit { + t.Fatal("guard not executed") + } +} + +func TestCheckDtoRequired(t *testing.T) { + f := New() + f.BindService(&TestSvc{}) + c := &Ctx{Ip: "x", MethodName: "Hello", ServiceName: "TestSvc"} + data := map[string]any{"age": 1} + c.Data = &data + + var streamCh chan any + var streamDone chan struct{} + _, err := f.invoke(c, &streamCh, &streamDone) + if err == nil { + t.Fatal("expected missing-field error") + } +} + +func TestGenCode(t *testing.T) { + GetFun().BindService(&TestSvc{}) + SetOutput(t.TempDir()) + GenCode(GenGo{}, GenTs{}) + if _, err := os.Stat(filepath.Join(getDirectory(), "go", "test_svc.go")); err != nil { + t.Fatalf("go service file missing: %v", err) + } + if _, err := os.Stat(filepath.Join(getDirectory(), "ts", "testSvc.ts")); err != nil { + t.Fatalf("ts service file missing: %v", err) + } +} + +func startServer(t *testing.T, port uint16) *Fun { + f := New() + f.BindService(&TestSvc{}) + go f.Start(port) + time.Sleep(300 * time.Millisecond) + return f +} + +// postCell 以标准库发起 /cell 调用并返回解码后的 Result +func postCell(t *testing.T, port uint16, body string) Result[any] { + t.Helper() + resp, err := http.Post(fmt.Sprintf("http://127.0.0.1:%d/cell", port), "application/json", + bytes.NewReader([]byte(body))) + if err != nil { + t.Fatal(err) + } + defer resp.Body.Close() + var out Result[any] + if err := json.NewDecoder(resp.Body).Decode(&out); err != nil { + t.Fatal(err) + } + return out +} + +func TestE2ERequest(t *testing.T) { + startServer(t, 39001) + res := postCell(t, 39001, `{"serviceName":"TestSvc","methodName":"Hello","data":{"name":"tom","age":1}}`) + if res.Status != 0 || res.Data == nil || (*res.Data).(string) != "hi tom" { + t.Fatalf("unexpected result: %+v", res) + } +} + +func TestE2EStream(t *testing.T) { + startServer(t, 39002) + resp, err := http.Post("http://127.0.0.1:39002/cell", "application/json", + bytes.NewReader([]byte(`{"serviceName":"TestSvc","methodName":"Count","data":{"name":"x","age":1}}`))) + if err != nil { + t.Fatal(err) + } + defer resp.Body.Close() + var got []string + scanner := bufio.NewScanner(resp.Body) + for scanner.Scan() { + line := bytes.TrimSpace(scanner.Bytes()) + if len(line) == 0 { + continue + } + got = append(got, string(line)) + } + if len(got) != 3 || got[0] != `"n0"` || got[2] != `"n2"` { + t.Fatalf("unexpected stream: %v", got) + } +} diff --git a/gen.go b/gen.go new file mode 100644 index 0000000..4a43916 --- /dev/null +++ b/gen.go @@ -0,0 +1,176 @@ +package fun + +import ( + "bytes" + "os" + "path/filepath" + "reflect" + "regexp" + "strings" + "text/template" +) + +type Gen interface { + typeToTemplateType(t reflect.Type) string + genService(svc *genSvc, serviceContext *genServiceType) + genDefaultService() + genStruct(t reflect.Type) *genImportType + getEnum(t reflect.Type) *genImportType + getName() string +} + +// genSvc 生成器视图下的服务 +type genSvc struct { + name string + methods []*genMethod +} + +// genMethod 生成器视图下的方法 +type genMethod struct { + name string + sig reflect.Type // 方法签名类型(不含接收者),Out(0) 为返回类型 + dtoType reflect.Type + isStream bool +} + +// serviceGroups 按服务名分组已注册方法 +func (f *Fun) serviceGroups() map[string][]*genMethod { + groups := map[string][]*genMethod{} + for key, m := range f.methods { + parts := strings.SplitN(key, ".", 2) + svc, name := parts[0], parts[1] + sig := reflect.New(m.serviceType).Method(m.methodIndex).Type() + groups[svc] = append(groups[svc], &genMethod{ + name: name, + sig: sig, + dtoType: m.dtoType, + isStream: m.isStream, + }) + } + return groups +} + +type genType struct { + GenServiceList []*genServiceType +} + +type genMethodType struct { + MethodName string + ReturnValueText string + DtoText string + ArgsText string + GenericTypeText string + IsProxy bool + IsStream bool +} + +type genEnumType struct { + Names []string + DisplayNames []string + Name string +} + +type genImportType struct { + Name string +} + +type genServiceType struct { + ServiceName string + GenMethodTypeList []*genMethodType + GenImport []*genImportType + IsIncludeProxy bool + IsIncludeStream bool +} + +type genClassType struct { + Name string + GenImport []*genImportType + GenClassFieldType []*genClassFieldType +} + +type genClassFieldType struct { + Name string + Type string + Tag string +} + +func deduplicateServiceImports(imports []*genImportType) []*genImportType { + seen := make(map[string]bool) + var result []*genImportType + for _, imp := range imports { + if !seen[imp.Name] { + seen[imp.Name] = true + result = append(result, imp) + } + } + return result +} + +func parseGenericTypeParams(typeName string) string { + start := strings.Index(typeName, "[") + end := strings.LastIndex(typeName, "]") + paramsStr := typeName[start+1 : end] + params := strings.Split(paramsStr, ",") + for i, param := range params { + LL := strings.Split(strings.TrimSpace(param), ".") + params[i] = firstLetterToUpper(LL[len(LL)-1]) + } + return strings.Join(params, "") +} + +func getGenericTypeName(typeName string) string { + start := strings.Index(typeName, "[") + return typeName[0:start] +} + +func genCode(templateContent string, outputFileName string, templateData any, languageName string) { + tmpl, err := template.New(languageName).Parse(templateContent) + if err != nil { + panic(err.Error()) + } + var buf bytes.Buffer + err = tmpl.Execute(&buf, templateData) + if err != nil { + panic(err.Error()) + } + code := buf.Bytes() + fullPath := filepath.Join(getDirectory(), languageName) + + _, err = os.Stat(fullPath) + if os.IsNotExist(err) { + err = os.MkdirAll(fullPath, os.ModePerm) + if err != nil { + panic(err.Error()) + } + } + err = os.WriteFile(filepath.Join(fullPath, outputFileName+"."+languageName), code, 0644) + if err != nil { + panic(err.Error()) + } +} + +// GenCode 执行代码生成:清空输出目录后按顺序运行每个生成器 +func GenCode(genList ...Gen) { + if err := os.RemoveAll(getDirectory()); err != nil && !os.IsNotExist(err) { + panic(err.Error()) + } + for _, gen := range genList { + gen.genDefaultService() + } +} + +var directory = "./gen" + +func SetOutput(path string) { + directory = path +} + +func getDirectory() string { + return directory +} + +func camelToSnake(s string) string { + re := regexp.MustCompile(`([a-z0-9])([A-Z])`) + snake := re.ReplaceAllString(s, `${1}_${2}`) + return strings.ToLower(snake) +} diff --git a/gen_go.go b/gen_go.go new file mode 100644 index 0000000..82bcd17 --- /dev/null +++ b/gen_go.go @@ -0,0 +1,222 @@ +package fun + +import ( + "reflect" + "strings" +) + +type GenGo struct { + template templateGo +} + +func (ctx GenGo) typeToTemplateType(t reflect.Type) string { + text := "" + if t.Kind() == reflect.Ptr { + text += "*" + t = t.Elem() + } + switch t.Kind() { + case reflect.Slice: + text += "[]" + ctx.typeToTemplateType(t.Elem()) + default: + text += t.Name() + } + return text +} + +func (ctx GenGo) genService(svc *genSvc, serviceContext *genServiceType) { + for _, gm := range svc.methods { + var returnValueText string + var dtoText string + var argsText string + var genericTypeText string + + // () error:无数据返回,客户端类型为 Void + if gm.sig.NumOut() == 1 && gm.sig.Out(0) == errorType { + genericTypeText = "Void" + returnValueText = "Result[Void]" + if gm.dtoType != nil { + v := ctx.typeToTemplateType(gm.dtoType) + if !strings.Contains(v, "[]") && strings.Contains(v, "[") { + dtoText += "dto " + getGenericTypeName(v) + parseGenericTypeParams(v) + } else { + dtoText += "dto " + v + } + argsText += ",dto" + ctx.genStruct(gm.dtoType) + } + serviceContext.GenMethodTypeList = append(serviceContext.GenMethodTypeList, &genMethodType{ + MethodName: gm.name, + ReturnValueText: returnValueText, + DtoText: dtoText, + ArgsText: argsText, + GenericTypeText: genericTypeText, + }) + continue + } + + returnType := gm.sig.Out(0) + if gm.isStream { + serviceContext.IsIncludeStream = true + if returnType == streamType { + genericTypeText = "any" + returnValueText = "Void" + } else { + t := ctx.typeToTemplateType(returnType) + if !strings.Contains(t, "[]") && strings.Contains(t, "[") { + genericTypeText = getGenericTypeName(t) + parseGenericTypeParams(t) + } else { + genericTypeText = t + } + returnValueText = genericTypeText + ctx.genReturnTypes(returnType) + } + } else { + t := ctx.typeToTemplateType(returnType) + if !strings.Contains(t, "[]") && strings.Contains(t, "[") { + returnValueText = getGenericTypeName(t) + parseGenericTypeParams(t) + } else { + returnValueText = t + } + genericTypeText = returnValueText + ctx.genReturnTypes(returnType) + returnValueText = "Result[" + returnValueText + "]" + } + + if gm.dtoType != nil { + v := ctx.typeToTemplateType(gm.dtoType) + if !strings.Contains(v, "[]") && strings.Contains(v, "[") { + dtoText += "dto " + getGenericTypeName(v) + parseGenericTypeParams(v) + } else { + dtoText += "dto " + v + } + argsText += ",dto" + ctx.genStruct(gm.dtoType) + } + + serviceContext.GenMethodTypeList = append(serviceContext.GenMethodTypeList, &genMethodType{ + MethodName: gm.name, + ReturnValueText: returnValueText, + DtoText: dtoText, + ArgsText: argsText, + GenericTypeText: genericTypeText, + IsStream: gm.isStream, + }) + } + genCode(ctx.template.genServiceTemplate(), camelToSnake(svc.name), serviceContext, ctx.getName()) +} + +// genReturnTypes 递归生成返回类型涉及的 struct/enum 定义 +func (ctx GenGo) genReturnTypes(returnType reflect.Type) { + if returnType.Kind() == reflect.Ptr { + returnType = returnType.Elem() + } + if returnType.Kind() == reflect.Struct { + ctx.genStruct(returnType) + } + if returnType.Kind() == reflect.Slice { + fieldType := returnType.Elem() + if fieldType.Kind() == reflect.Ptr { + fieldType = fieldType.Elem() + } + if fieldType.Kind() == reflect.Struct { + ctx.genStruct(fieldType) + } + } + if returnType.Kind() == reflect.Uint8 && (returnType.Implements(displayEnumType) || returnType.Implements(enumType)) { + ctx.getEnum(returnType) + } +} + +func (ctx GenGo) genDefaultService() { + f := GetFun() + genContext := genType{GenServiceList: []*genServiceType{}} + + for svcName, methods := range f.serviceGroups() { + serviceContext := &genServiceType{ + ServiceName: svcName, + GenMethodTypeList: []*genMethodType{}, + } + genContext.GenServiceList = append(genContext.GenServiceList, serviceContext) + ctx.genService(&genSvc{name: svcName, methods: methods}, serviceContext) + } + genCode(ctx.template.genDefaultServiceTemplate(), "fun", genContext, ctx.getName()) +} + +func (ctx GenGo) genStruct(t reflect.Type) *genImportType { + var structTemplate genClassType + if !strings.Contains(t.String(), "[]") && strings.Contains(t.String(), "[") { + structTemplate = genClassType{ + Name: getGenericTypeName(t.Name()) + parseGenericTypeParams(t.Name()), + } + } else { + structTemplate = genClassType{ + Name: t.Name(), + } + } + + for i := 0; i < t.NumField(); i++ { + field := t.Field(i) + fieldType := field.Type + jsType := ctx.typeToTemplateType(fieldType) + name := field.Name + tag := "`json:\"" + firstLetterToLower(name) + "\"`" + if !strings.Contains(jsType, "[]") && strings.Contains(jsType, "[") { + structTemplate.GenClassFieldType = append(structTemplate.GenClassFieldType, &genClassFieldType{ + Name: name, + Type: getGenericTypeName(jsType) + parseGenericTypeParams(jsType), + Tag: tag, + }) + } else { + structTemplate.GenClassFieldType = append(structTemplate.GenClassFieldType, &genClassFieldType{ + Name: name, + Type: jsType, + Tag: tag, + }) + } + + if fieldType.Kind() == reflect.Struct { + ctx.genStruct(fieldType) + } + if fieldType.Kind() == reflect.Slice && fieldType.Elem().Kind() == reflect.Struct { + ctx.genStruct(fieldType.Elem()) + } + if fieldType.Kind() == reflect.Uint8 && (fieldType.Implements(displayEnumType) || fieldType.Implements(enumType)) { + ctx.getEnum(fieldType) + } + } + + genCode( + ctx.template.genStructTemplate(), + camelToSnake(structTemplate.Name), + structTemplate, + ctx.getName(), + ) + return &genImportType{} +} + +func (ctx GenGo) getEnum(t reflect.Type) *genImportType { + var enumTemplate genEnumType + statusValue := reflect.New(t).Elem() + if t.Implements(displayEnumType) { + enumValue := statusValue.Interface().(displayEnum) + enumTemplate.Names = enumValue.Names() + enumTemplate.DisplayNames = enumValue.DisplayNames() + } else { + enumValue := statusValue.Interface().(enum) + enumTemplate.Names = enumValue.Names() + } + enumTemplate.Name = t.Name() + + genCode( + ctx.template.genEnumTemplate(), + camelToSnake(t.Name()), + enumTemplate, + ctx.getName(), + ) + return &genImportType{} +} + +func (ctx GenGo) getName() string { + return "go" +} diff --git a/gen_ts.go b/gen_ts.go new file mode 100644 index 0000000..4e66a7b --- /dev/null +++ b/gen_ts.go @@ -0,0 +1,250 @@ +package fun + +import ( + "reflect" + "strings" +) + +type GenTs struct { + template templateTs +} + +func (ctx GenTs) typeToTemplateType(t reflect.Type) string { + text := "" + if t.Kind() == reflect.Ptr { + text += " | null" + t = t.Elem() + } + switch t.Kind() { + case reflect.Int8, reflect.Int16, reflect.Int32, reflect.Int64, + reflect.Uint8, reflect.Uint16, reflect.Uint32, reflect.Uint64: + if t.Kind() == reflect.Uint8 && (t.Implements(displayEnumType) || t.Implements(enumType)) { + text = t.Name() + text + } else { + text = "number" + text + } + case reflect.Bool: + text = "boolean" + text + case reflect.String, reflect.Struct: + text = t.Name() + text + default: + text = ctx.typeToTemplateType(t.Elem()) + "[]" + text + } + return text +} + +func (ctx GenTs) genService(svc *genSvc, serviceContext *genServiceType) { + var nestedImports []*genImportType + + for _, gm := range svc.methods { + var returnValueText string + var dtoText string + var argsText string + var genericTypeText string + + // () error:无数据返回,客户端类型为 void + if gm.sig.NumOut() == 1 && gm.sig.Out(0) == errorType { + genericTypeText = "void" + returnValueText = "result" + if gm.dtoType != nil { + v := firstLetterToLower(ctx.typeToTemplateType(gm.dtoType)) + if !strings.Contains(v, "[]") && strings.Contains(v, "[") { + dtoText += "dto:" + getGenericTypeName(v) + parseGenericTypeParams(v) + } else { + dtoText += "dto:" + v + } + argsText += ",dto" + nestedImports = append(nestedImports, ctx.genStruct(gm.dtoType)) + } + serviceContext.GenMethodTypeList = append(serviceContext.GenMethodTypeList, &genMethodType{ + MethodName: firstLetterToLower(gm.name), + ReturnValueText: returnValueText, + DtoText: dtoText, + ArgsText: argsText, + GenericTypeText: genericTypeText, + }) + continue + } + + returnType := gm.sig.Out(0) + if gm.isStream { + serviceContext.IsIncludeStream = true + if returnType == streamType { + genericTypeText = "any" + returnValueText = "void" + } else { + t := firstLetterToLower(ctx.typeToTemplateType(returnType)) + if !strings.Contains(t, "[]") && strings.Contains(t, "[") { + genericTypeText = getGenericTypeName(t) + parseGenericTypeParams(t) + } else { + genericTypeText = t + } + returnValueText = "void" + nestedImports = ctx.genReturnTypes(returnType, nestedImports) + } + } else { + t := firstLetterToLower(ctx.typeToTemplateType(returnType)) + if !strings.Contains(t, "[]") && strings.Contains(t, "[") { + returnValueText = getGenericTypeName(t) + parseGenericTypeParams(t) + } else { + returnValueText = t + } + genericTypeText = returnValueText + nestedImports = ctx.genReturnTypes(returnType, nestedImports) + returnValueText = "result<" + returnValueText + ">" + } + + if gm.dtoType != nil { + v := firstLetterToLower(ctx.typeToTemplateType(gm.dtoType)) + if !strings.Contains(v, "[]") && strings.Contains(v, "[") { + dtoText += "dto:" + getGenericTypeName(v) + parseGenericTypeParams(v) + } else { + dtoText += "dto:" + v + } + argsText += ",dto" + nestedImports = append(nestedImports, ctx.genStruct(gm.dtoType)) + } + + serviceContext.GenMethodTypeList = append(serviceContext.GenMethodTypeList, &genMethodType{ + MethodName: firstLetterToLower(gm.name), + ReturnValueText: returnValueText, + DtoText: dtoText, + ArgsText: argsText, + GenericTypeText: firstLetterToLower(genericTypeText), + IsStream: gm.isStream, + }) + } + serviceContext.GenImport = deduplicateServiceImports(nestedImports) + + genCode( + ctx.template.genServiceTemplate(), + firstLetterToLower(svc.name), + serviceContext, + ctx.getName(), + ) +} + +// genReturnTypes 递归生成返回类型涉及的 struct/enum 导入 +func (ctx GenTs) genReturnTypes(returnType reflect.Type, nestedImports []*genImportType) []*genImportType { + if returnType.Kind() == reflect.Ptr { + returnType = returnType.Elem() + } + if returnType.Kind() == reflect.Struct { + nestedImports = append(nestedImports, ctx.genStruct(returnType)) + } + if returnType.Kind() == reflect.Slice { + fieldType := returnType.Elem() + if fieldType.Kind() == reflect.Ptr { + fieldType = fieldType.Elem() + } + if fieldType.Kind() == reflect.Struct { + nestedImports = append(nestedImports, ctx.genStruct(fieldType)) + } + } + if returnType.Kind() == reflect.Uint8 && (returnType.Implements(displayEnumType) || returnType.Implements(enumType)) { + nestedImports = append(nestedImports, ctx.getEnum(returnType)) + } + return nestedImports +} + +func (ctx GenTs) genDefaultService() { + f := GetFun() + genContext := genType{GenServiceList: []*genServiceType{}} + + for svcName, methods := range f.serviceGroups() { + serviceContext := &genServiceType{ + ServiceName: firstLetterToLower(svcName), + GenMethodTypeList: []*genMethodType{}, + } + genContext.GenServiceList = append(genContext.GenServiceList, serviceContext) + ctx.genService(&genSvc{name: svcName, methods: methods}, serviceContext) + } + genCode(ctx.template.genClientTemplate(), "client", nil, ctx.getName()) + genCode(ctx.template.genDefaultServiceTemplate(), "fun", genContext, ctx.getName()) +} + +func (ctx GenTs) genStruct(t reflect.Type) *genImportType { + var structTemplate genClassType + if !strings.Contains(t.String(), "[]") && strings.Contains(t.String(), "[") { + structTemplate = genClassType{ + Name: firstLetterToLower(getGenericTypeName(t.Name())) + parseGenericTypeParams(t.Name()), + } + } else { + structTemplate = genClassType{ + Name: firstLetterToLower(t.Name()), + } + } + var nestedImports []*genImportType + + for i := 0; i < t.NumField(); i++ { + field := t.Field(i) + fieldType := field.Type + jsType := ctx.typeToTemplateType(fieldType) + name := field.Name + if fieldType.Kind() == reflect.Ptr { + fieldType = fieldType.Elem() + name += "?" + } + if !strings.Contains(jsType, "[]") && strings.Contains(jsType, "[") { + structTemplate.GenClassFieldType = append(structTemplate.GenClassFieldType, &genClassFieldType{ + Name: firstLetterToLower(name), + Type: firstLetterToLower(getGenericTypeName(jsType)) + parseGenericTypeParams(jsType), + }) + } else { + structTemplate.GenClassFieldType = append(structTemplate.GenClassFieldType, &genClassFieldType{ + Name: firstLetterToLower(name), + Type: firstLetterToLower(jsType), + }) + } + + if fieldType.Kind() == reflect.Struct { + nestedImports = append(nestedImports, ctx.genStruct(fieldType)) + } + if fieldType.Kind() == reflect.Slice && fieldType.Elem().Kind() == reflect.Struct { + nestedImports = append(nestedImports, ctx.genStruct(fieldType.Elem())) + } + if fieldType.Kind() == reflect.Uint8 && (fieldType.Implements(displayEnumType) || fieldType.Implements(enumType)) { + nestedImports = append(nestedImports, ctx.getEnum(fieldType)) + } + } + + structTemplate.GenImport = deduplicateServiceImports(nestedImports) + + genCode( + ctx.template.genStructTemplate(), + structTemplate.Name, + structTemplate, + ctx.getName(), + ) + + if !strings.Contains(t.String(), "[]") && strings.Contains(t.String(), "[") { + return &genImportType{Name: structTemplate.Name} + } + return &genImportType{Name: firstLetterToLower(t.Name())} +} + +func (ctx GenTs) getEnum(t reflect.Type) *genImportType { + var enumTemplate genEnumType + statusValue := reflect.New(t).Elem() + if t.Implements(displayEnumType) { + enumValue := statusValue.Interface().(displayEnum) + enumTemplate.Names = enumValue.Names() + enumTemplate.DisplayNames = enumValue.DisplayNames() + } else { + enumValue := statusValue.Interface().(enum) + enumTemplate.Names = enumValue.Names() + } + enumTemplate.Name = firstLetterToLower(t.Name()) + + genCode( + ctx.template.genEnumTemplate(), + firstLetterToLower(t.Name()), + enumTemplate, + ctx.getName(), + ) + return &genImportType{Name: firstLetterToLower(t.Name())} +} + +func (ctx GenTs) getName() string { + return "ts" +} diff --git a/go.mod b/go.mod new file mode 100644 index 0000000..c1bf427 --- /dev/null +++ b/go.mod @@ -0,0 +1,11 @@ +module github.com/cyi-cc/fun + +go 1.26.1 + +require github.com/valyala/fasthttp v1.73.0 + +require ( + github.com/andybalholm/brotli v1.2.2 // indirect + github.com/klauspost/compress v1.19.1 // indirect + github.com/valyala/bytebufferpool v1.0.0 // indirect +) diff --git a/go.sum b/go.sum new file mode 100644 index 0000000..32ddbf8 --- /dev/null +++ b/go.sum @@ -0,0 +1,10 @@ +github.com/andybalholm/brotli v1.2.2 h1:HzTuoo2ErYQqf5qvcJInB8uvqSVxRttzkFexPWtnceM= +github.com/andybalholm/brotli v1.2.2/go.mod h1:rzTDkvFWvIrjDXZHkuS16NPggd91W3kUSvPlQ1pLaKY= +github.com/klauspost/compress v1.19.1 h1:VsB4HPswih7mmZ8WleSFQ75c/Ui1M4trX5oAsJnhSlk= +github.com/klauspost/compress v1.19.1/go.mod h1:cwPg85FWrGar70rWktvGQj8/hthj3wpl0PGDogxkrSQ= +github.com/valyala/bytebufferpool v1.0.0 h1:GqA5TC/0021Y/b9FG4Oi9Mr3q7XYx6KllzawFIhcdPw= +github.com/valyala/bytebufferpool v1.0.0/go.mod h1:6bBcMArwyJ5K/AmCkWv1jt77kVWyCJ6HpOuEn7z0Csc= +github.com/valyala/fasthttp v1.73.0 h1:ocTOORnBWtJ+P8t/6wAjdkchMzdfHmWx2VD/DPbgZ7s= +github.com/valyala/fasthttp v1.73.0/go.mod h1:EtXQDHaR+5P18p8wqDRFpUhxr108Ga9mXvVJXHRrN2k= +github.com/xyproto/randomstring v1.0.5 h1:YtlWPoRdgMu3NZtP45drfy1GKoojuR7hmRcnhZqKjWU= +github.com/xyproto/randomstring v1.0.5/go.mod h1:rgmS5DeNXLivK7YprL0pY+lTuhNQW3iGxZ18UQApw/E= diff --git a/guard.go b/guard.go new file mode 100644 index 0000000..54669ca --- /dev/null +++ b/guard.go @@ -0,0 +1,5 @@ +package fun + +type Guard interface { + Guard(ctx Ctx) +} diff --git a/handle.go b/handle.go new file mode 100644 index 0000000..df1dc53 --- /dev/null +++ b/handle.go @@ -0,0 +1,226 @@ +package fun + +import ( + "bufio" + "encoding/json" + "errors" + "fmt" + "reflect" + "runtime/debug" + + "github.com/valyala/fasthttp" +) + +// handle 处理 HTTP 请求 +func (f *Fun) handle(fastCtx *fasthttp.RequestCtx) { + ctx := &Ctx{RequestCtx: fastCtx} + defer f.handlePanic(ctx) + + if ctx.path() != "/cell" { + ctx.setStatusCode(fasthttp.StatusNotFound) + return + } + if !ctx.isPost() { + ctx.setStatusCode(fasthttp.StatusMethodNotAllowed) + return + } + + body := ctx.postBody() + var requestInfo RequestInfo[map[string]any] + if err := json.Unmarshal(body, &requestInfo); err != nil { + ctx.sendError(err) + return + } + requestInfo.MethodName = firstLetterToUpper(requestInfo.MethodName) + requestInfo.ServiceName = firstLetterToUpper(requestInfo.ServiceName) + if requestInfo.MethodName == "" || requestInfo.ServiceName == "" { + ctx.sendError(errEmptyFields) + return + } + + ctx.Ip = ctx.remoteIP().String() + ctx.State = requestInfo.State + ctx.MethodName = requestInfo.MethodName + ctx.ServiceName = requestInfo.ServiceName + ctx.Data = requestInfo.Data + + // 流式方法:响应保持打开,以 NDJSON 行推送(Streamable HTTP) + // streamCh != nil 表示流式方法;业务返回的 *Stream 在 invoke 内完成注入 + var streamCh chan any + var streamDone chan struct{} + + result, err := f.invoke(ctx, &streamCh, &streamDone) + if err != nil { + ctx.sendError(err) + return + } + if streamCh != nil { + fastCtx.Response.Header.SetContentType("application/x-ndjson") + fastCtx.Response.Header.Set("Cache-Control", "no-cache") + fastCtx.Response.Header.Set("Connection", "keep-alive") + fastCtx.SetBodyStreamWriter(func(w *bufio.Writer) { + writeLine := func(v any) bool { + data, err := json.Marshal(v) + if err != nil { + return false + } + raw, err := lowerKeysFromJSON(data) + if err == nil { + if data, err = json.Marshal(raw); err != nil { + return false + } + } + if _, err := fmt.Fprintf(w, "%s\n", data); err != nil { + return false + } + if err := w.Flush(); err != nil { + return false + } + return true + } + // (T, stream, error):T 作为流的第一条消息下发 + if result.Data != nil { + if !writeLine(*result.Data) { + close(streamDone) + return + } + } + for message := range streamCh { + if !writeLine(message) { + close(streamDone) + return + } + } + close(streamDone) + }) + return + } + ctx.send(*result) +} + +// handlePanic 兜底处理 panic:归一为 error 后写回错误响应,并记录完整堆栈日志 +func (f *Fun) handlePanic(c *Ctx) { + if v := recover(); v != nil { + var err error + if e, ok := v.(error); ok { + err = e + } else { + err = fmt.Errorf("panic (%s.%s): %v", c.ServiceName, c.MethodName, v) + } + ErrorLogger(err.Error(), "\n"+string(debug.Stack())) + c.sendError(err) + } +} + +// invoke 按 "Service.Method" 查找并调用,返回成功结果 +// 流式方法时,业务返回的 *Stream 完成通道注入(streamCh/streamDone 被创建并填充) +// 预期错误(方法不存在、参数缺失、业务失败)以 error 返回,不 panic +func (f *Fun) invoke(c *Ctx, streamCh *chan any, streamDone *chan struct{}) (*Result[any], error) { + key := c.ServiceName + "." + c.MethodName + method, ok := f.methods[key] + if !ok { + return nil, errMethodNotFound + } + + f.callGuard(c, c.ServiceName) + + var args []reflect.Value + if method.dtoType != nil { + if c.Data == nil { + return nil, errDTORequired + } + if err := checkDto(method.dtoType, *c.Data, c.MethodName); err != nil { + return nil, err + } + dto := reflect.New(method.dtoType).Elem() + if err := convert(c.Data, dto.Addr().Interface()); err != nil { + return nil, err + } + args = append(args, dto) + } + + // 每请求创建新实例并注入 Ctx/Box 依赖,避免并发共享实例 + instance := reflect.New(method.serviceType) + f.serviceWired(instance.Elem(), c) + values := instance.Method(method.methodIndex).Call(args) + return callResult(c, values, method, streamCh, streamDone) +} + +// callResult 将反射调用结果归一为 Result +// 兼容四种签名:(error)、(T, error)、(stream, error)、(T, stream, error) +// - 末位返回值是 error 且非 nil → 业务失败,返回 error +// - 带 *Stream 的签名:注入推送通道后,仅 (T, stream, error) 返回 T 作为数据 +func callResult(c *Ctx, values []reflect.Value, method methodInfo, streamCh *chan any, streamDone *chan struct{}) (*Result[any], error) { + if last := values[len(values)-1]; last.Type().Implements(errorType) { + if !last.IsNil() { + // 业务出错但可能已启动 goroutine 调 Send/Close: + // 注入一个已取消的流,让它们立即解除阻塞退出,避免 goroutine 泄漏 + if method.isStream { + injectCancelledStream(values, method, streamCh, streamDone) + } + // 业务 Error() 构造的 Result[any] 作为 error 返回,sendError 里原样透传 + var result Result[any] + if errors.As(last.Interface().(error), &result) { + return nil, result + } + return nil, last.Interface().(error) + } + values = values[:len(values)-1] + } + + if method.isStream { + // (stream, error):流在第 0 位;(T, stream, error):流在第 1 位 + streamIdx := 0 + if len(values) == 2 { + streamIdx = 1 + } + s := values[streamIdx].Interface().(*Stream) + if s == nil { + return nil, errors.New("fun: method returned nil stream") + } + *streamCh = make(chan any) + *streamDone = make(chan struct{}) + s.Inject(*streamCh, *streamDone) + // (T, stream, error):返回 T;纯流:不返回数据 + if len(values) == 2 { + r := success(values[0].Interface()) + return &r, nil + } + r := success(nil) + return &r, nil + } + + if len(values) == 0 { + // () error:无数据返回 + r := success(nil) + return &r, nil + } + r := success(values[0].Interface()) + return &r, nil +} + +// injectCancelledStream 业务出错时注入已取消的流通道, +// 使正在 Send/Close 上阻塞的业务 goroutine 立即解除并退出 +func injectCancelledStream(values []reflect.Value, method methodInfo, streamCh *chan any, streamDone *chan struct{}) { + streamIdx := 0 + if len(values) == 3 { // (T, stream, error) + streamIdx = 1 + } + s, ok := values[streamIdx].Interface().(*Stream) + if !ok || s == nil { + return + } + *streamCh = make(chan any) + *streamDone = make(chan struct{}) + s.Inject(*streamCh, *streamDone) + close(*streamDone) +} + +// convert 将数据转为 JSON 再反序列化到目标类型,避免手写字段映射 +func convert(from any, to any) error { + data, err := json.Marshal(from) + if err != nil { + return err + } + return json.Unmarshal(data, to) +} diff --git a/logger.go b/logger.go new file mode 100644 index 0000000..f0a39cb --- /dev/null +++ b/logger.go @@ -0,0 +1,384 @@ +package fun + +import ( + "bytes" + "encoding/json" + "fmt" + "os" + "path/filepath" + "runtime" + "sort" + "strconv" + "strings" + "sync" + "time" +) + +const ( + PanicLevel uint8 = iota + ErrorLevel + WarnLevel + InfoLevel + DebugLevel + TraceLevel +) + +var logChan = make(chan string, 100) +var logWg sync.WaitGroup + +const ( + TerminalMode uint8 = iota + FileMode +) + +var logMutex sync.Mutex + +type Logger struct { + Level uint8 + Mode uint8 + MaxSizeFile uint8 //文件最大大小(MB) + MaxNumberFiles uint64 //文件最多数量 + ExpireLogsDays uint8 //文件保留时间 + LogFilePath string +} + +var logger Logger = Logger{ + Level: TraceLevel, + Mode: TerminalMode, + MaxSizeFile: 0, + MaxNumberFiles: 0, + ExpireLogsDays: 0, + LogFilePath: "../log", +} + +func init() { + go deleteLogWorker() + go logWriterWorker() +} + +func logWriterWorker() { + for text := range logChan { + logMutex.Lock() + if logger.Mode == FileMode { + fileLogger(text) + } else { + fmt.Println(text) + } + } +} + +func deleteLogWorker() { + cleanupExpiredLogs() + ticker := time.NewTicker(24 * time.Hour) + defer ticker.Stop() + for { + select { + case <-ticker.C: + if logger.Mode == FileMode { + cleanupExpiredLogs() + } + } + } +} + +func getLogFilePath() string { + if logger.LogFilePath == "" { + return "./log" + } + return logger.LogFilePath +} + +func cleanupExpiredLogs() { + if logger.ExpireLogsDays <= 0 { + return + } + _, err := os.Stat(getLogFilePath()) + if os.IsNotExist(err) { + return + } + if err != nil { + return + } + entries, err := os.ReadDir(getLogFilePath()) + if err != nil { + return + } + expireDuration := time.Duration(logger.ExpireLogsDays) * 24 * time.Hour + currentTimeMillis := time.Now().UnixMilli() + expireThreshold := currentTimeMillis - expireDuration.Milliseconds() + + for _, entry := range entries { + if !entry.IsDir() { + fileNameInfo := getFileNameInfo(entry.Name()) + if fileNameInfo.LoggerTime == 0 { + continue + } + if fileNameInfo.LoggerTime < expireThreshold { + fullPath := filepath.Join(getLogFilePath(), entry.Name()) + err := os.Remove(fullPath) + if err != nil && !os.IsNotExist(err) { + return + } + } + } + } +} + +func getFileNameInfo(name string) fileName { + fileNameParts := strings.Split(name, ".log.") + if len(fileNameParts) != 2 { + deleteLog(name) + return fileName{} + } + dateLayout := "2006-01-02" + dateString := fileNameParts[0] + fileDate, err := time.Parse(dateLayout, dateString) + if err != nil { + deleteLog(name) + return fileName{} + } + indexString := fileNameParts[1] + indexString = strings.TrimSuffix(indexString, ".log") + fileIndex, err := strconv.ParseInt(indexString, 10, 32) + if err != nil { + deleteLog(name) + return fileName{} + } + return fileName{ + index: int32(fileIndex), + LoggerTime: fileDate.UnixMilli(), + } +} + +type fileName struct { + LoggerTime int64 + index int32 +} + +func deleteLog(name string) { + fullPath := filepath.Join(getLogFilePath(), name) + err := os.Remove(fullPath) + if err != nil && !os.IsNotExist(err) { + return + } +} + +func fileLogger(text string) { + _, err := os.Stat(getLogFilePath()) + if os.IsNotExist(err) { + err = os.MkdirAll(getLogFilePath(), os.ModePerm) + if err != nil { + return + } + } + currentDate := getCurrentData() + logFileName := currentDate + ".log" + logFilePath := filepath.Join(getLogFilePath(), logFileName) + logFilePath, err = getNextLogFile(getLogFilePath(), currentDate, text) + if err != nil { + return + } + file, err := os.OpenFile(logFilePath, os.O_CREATE|os.O_WRONLY|os.O_APPEND, 0644) + if err != nil { + return + } + defer func(file *os.File) { + _ = file.Close() + }(file) + _, _ = file.WriteString(text + "\n") +} + +func removeOldestLogFile(entries []os.DirEntry) { + if logger.MaxNumberFiles == 0 { + return + } + if uint64(len(entries)) < logger.MaxNumberFiles { + return + } + var newEntries []fileName + for _, v := range entries { + fileNameInfo := getFileNameInfo(v.Name()) + if fileNameInfo.LoggerTime != 0 { + newEntries = append(newEntries, fileNameInfo) + } + } + if uint64(len(newEntries)) < logger.MaxNumberFiles { + return + } + delNum := uint64(len(newEntries)) - logger.MaxNumberFiles + 1 + sort.Slice(newEntries, func(i, j int) bool { + if newEntries[i].LoggerTime != newEntries[j].LoggerTime { + return newEntries[i].LoggerTime < newEntries[j].LoggerTime + } + return newEntries[i].index < newEntries[j].index + }) + for i := 0; i < int(delNum); i++ { + fileName := newEntries[i] + t := time.Unix(0, fileName.LoggerTime*int64(time.Millisecond)) + fileNamePath := filepath.Join(getLogFilePath(), t.Format("2006-01-02")+".log."+strconv.Itoa(int(fileName.index))) + deleteLog(fileNamePath) + } +} + +// getNextLogFile 获取下一个应该写入的日志文件 +func getNextLogFile(dirPath, dateStr string, text string) (string, error) { + entries, err := os.ReadDir(dirPath) + if err != nil { + return filepath.Join(dirPath, dateStr+".log.1"), err + } + var maxIndex int32 = 0 + var existingFiles []string + for _, entry := range entries { + if !entry.IsDir() && strings.HasPrefix(entry.Name(), dateStr+".log") { + existingFiles = append(existingFiles, entry.Name()) + } + } + if len(existingFiles) == 0 { + removeOldestLogFile(entries) + return filepath.Join(dirPath, dateStr+".log.1"), nil + } + for _, fileName := range existingFiles { + fileNameInfo := getFileNameInfo(fileName) + if fileNameInfo.LoggerTime != 0 && fileNameInfo.index > maxIndex { + maxIndex = fileNameInfo.index + } + } + if maxIndex == 0 { + removeOldestLogFile(entries) + return filepath.Join(dirPath, dateStr+".log.1"), nil + } + if logger.MaxSizeFile > 0 && maxIndex > 0 { + currentFile := filepath.Join(dirPath, fmt.Sprintf("%s.log.%d", dateStr, maxIndex)) + if fileInfo, err := os.Stat(currentFile); err == nil { + maxSizeBytes := int64(logger.MaxSizeFile) * 1024 * 1024 + if fileInfo.Size()+int64(len(text)) > maxSizeBytes { + removeOldestLogFile(entries) + return filepath.Join(dirPath, fmt.Sprintf("%s.log.%d", dateStr, maxIndex+1)), nil + } + } else { + return "", err + } + } + return filepath.Join(dirPath, fmt.Sprintf("%s.log.%d", dateStr, maxIndex)), nil +} + +func ConfigLogger(log Logger) { + logger = log +} + +func getCurrentTime() string { + return time.Now().Format("2006-01-02 15:04:05") +} + +func getCurrentData() string { + return time.Now().Format("2006-01-02") +} + +func getMethodNameLogger() string { + pc, _, _, _ := runtime.Caller(3) + fn := runtime.FuncForPC(pc) + charsToRemove := []string{"(", "*", ")"} + name := fn.Name() + for _, char := range charsToRemove { + name = strings.ReplaceAll(name, char, "") + } + funcName := "[" + padString(strings.ReplaceAll(name, "/", "."), 40) + "] " + return funcName +} + +func getLevelName(level uint8) string { + switch level { + case TraceLevel: + return "TRACE" + case DebugLevel: + return "DEBUG" + case InfoLevel: + return "INFO" + case ErrorLevel: + return "ERROR" + case WarnLevel: + return "WARN" + default: + return "PANIC" + } +} + +func sendLogWorker(level uint8, message []any) { + if logger.Level >= level { + var text1 strings.Builder + for _, m := range message { + var msgStr string + var temp interface{} + var trimmedStr string + switch v := m.(type) { + case string: + err := json.Unmarshal([]byte(v), &temp) + if err != nil { + msgStr = fmt.Sprintf("%s", v) + break + } + bs, _ := json.Marshal(&temp) + trimmedStr = string(bs) + case []byte: + err := json.Unmarshal(v, &temp) + if err != nil { + msgStr = fmt.Sprintf("%s", v) + break + } + bs, _ := json.Marshal(&temp) + trimmedStr = string(bs) + default: + bs, _ := json.Marshal(v) + err := json.Unmarshal(bs, &temp) + if err != nil { + msgStr = fmt.Sprintf("%v", v) + break + } + trimmedStr = string(bs) + } + switch temp.(type) { + case map[string]any, []any: + var out bytes.Buffer + err := json.Indent(&out, []byte(trimmedStr), "", "\t") + if err != nil { + return + } + msgStr = fmt.Sprintf("\n%s", out.String()) + default: + msgStr = fmt.Sprintf("%v", m) + } + text1.WriteString(msgStr + " ") + } + text := "[" + getCurrentTime() + "] [" + padString(getLevelName(level), 7) + "] " + getMethodNameLogger() + text1.String() + logWg.Add(1) + logChan <- text + } +} + +func DebugLogger(message ...any) { + sendLogWorker(DebugLevel, message) +} + +func InfoLogger(message ...any) { + sendLogWorker(InfoLevel, message) +} + +func TraceLogger(message ...any) { + sendLogWorker(TraceLevel, message) +} + +func ErrorLogger(message ...any) { + sendLogWorker(ErrorLevel, message) +} + +func WarnLogger(message ...any) { + sendLogWorker(WarnLevel, message) +} + +func PanicLogger(message ...any) { + sendLogWorker(PanicLevel, message) +} + +func padString(str string, totalLength int) string { + return fmt.Sprintf("%-*s", totalLength, str)[0:totalLength] +} diff --git a/request.go b/request.go new file mode 100644 index 0000000..db9680c --- /dev/null +++ b/request.go @@ -0,0 +1,14 @@ +package fun + +const ( + RequestNormalType uint8 = iota + RequestStreamType +) + +type RequestInfo[T any] struct { + MethodName string + ServiceName string + Data *T + State map[string]string + Type uint8 +} diff --git a/result.go b/result.go new file mode 100644 index 0000000..765d41a --- /dev/null +++ b/result.go @@ -0,0 +1,60 @@ +package fun + +import ( + "fmt" + "reflect" +) + +const ( + successCode uint8 = iota + cellErrorCode + errorCode +) + +type Result[T any] struct { + Id string `json:"id,omitempty"` + Code *uint16 `json:"code,omitempty"` + Data *T `json:"data,omitempty"` + Msg *string `json:"msg,omitempty"` + Status uint8 `json:"status"` +} + +// Error 让 Result 实现 error 接口,业务方法可直接返回,Code/Msg/Status 随结果透传 +func (r Result[T]) Error() string { + if r.Msg != nil { + return *r.Msg + } + if r.Code != nil { + return fmt.Sprintf("code=%d", *r.Code) + } + return "fun: unknown error" +} + +// Error 构造带错误码的错误响应,作为 error 返回 +// 用法:return "", fun.Error(4001, "登录失败") +func Error(code uint16, msg string) error { + return Result[any]{Code: &code, Msg: &msg, Status: errorCode} +} + +func callError(err error) Result[any] { + return Result[any]{Msg: new(err.Error()), Status: cellErrorCode} +} + +// success 构造成功响应,空切片规范化为 [] 而不是 null +func success(data any) Result[any] { + return Result[any]{Data: nonNil(data), Status: successCode} +} + +// nonNil 返回 data 的指针;空切片会重建为同类型的非 nil 空切片, +// 保证 JSON 序列化输出 [] 而不是 null +func nonNil(data any) *any { + if data == nil { + return nil + } + + v := reflect.ValueOf(data) + if v.Kind() == reflect.Slice && v.Len() == 0 { + return new(reflect.MakeSlice(v.Type(), 0, 0).Interface()) + } + return &data +} diff --git a/stream.go b/stream.go new file mode 100644 index 0000000..afe9199 --- /dev/null +++ b/stream.go @@ -0,0 +1,80 @@ +package fun + +import ( + "fmt" + "sync" +) + +// Stream 流式响应的业务句柄。 +// 业务方法返回 *Stream 后,框架调用 Inject 注入推送通道; +// 未注入时 Send/Close 自动阻塞等待,避免业务 goroutine 与注入之间的竞态。 +type Stream struct { + mu sync.Mutex + once sync.Once + ready chan struct{} + ch chan any + done chan struct{} + closed bool + onClose func() +} + +func (s *Stream) getReady() chan struct{} { + s.once.Do(func() { + if s.ready == nil { + s.ready = make(chan struct{}) + } + }) + return s.ready +} + +// Inject 注入推送通道与结束信号,由框架在方法返回后调用 +func (s *Stream) Inject(ch chan any, done chan struct{}) { + s.ch = ch + s.done = done + close(s.getReady()) +} + +// Send 推送一条消息;连接断开或流已关闭时返回错误 +func (s *Stream) Send(message any) error { + <-s.getReady() + s.mu.Lock() + defer s.mu.Unlock() + if s.closed { + return fmt.Errorf("fun: stream closed") + } + select { + case s.ch <- message: + return nil + case <-s.done: + return fmt.Errorf("fun: stream closed") + } +} + +// Close 主动结束流,触发 OnClose 回调 +func (s *Stream) Close() { + <-s.getReady() + s.mu.Lock() + if s.closed { + s.mu.Unlock() + return + } + s.closed = true + cb := s.onClose + s.mu.Unlock() + if cb != nil { + cb() + } + close(s.ch) +} + +// OnClose 注册关闭回调;流已关闭时立即执行 +func (s *Stream) OnClose(cb func()) { + <-s.getReady() + s.mu.Lock() + defer s.mu.Unlock() + if s.closed { + cb() + return + } + s.onClose = cb +} diff --git a/tag.go b/tag.go new file mode 100644 index 0000000..00f9e89 --- /dev/null +++ b/tag.go @@ -0,0 +1,34 @@ +package fun + +import ( + "reflect" + "strings" +) + +type Tag struct { + TagList map[string]string +} + +func newTag(tag reflect.StructTag) *Tag { + t := &Tag{ + TagList: map[string]string{}, + } + pairs := strings.Split(strings.TrimSpace(tag.Get("fun")), ";") + for _, pair := range pairs { + if pair == "" { + continue + } + keyValue := strings.Split(pair, ":") + if len(keyValue) == 1 { + t.TagList[keyValue[0]] = "" + } else { + t.TagList[keyValue[0]] = keyValue[1] + } + } + return t +} + +func (tag *Tag) getTag(key string) (string, bool) { + v, ok := tag.TagList[key] + return v, ok +} diff --git a/template_go.go b/template_go.go new file mode 100644 index 0000000..561f110 --- /dev/null +++ b/template_go.go @@ -0,0 +1,278 @@ +package fun + +type templateGo struct{} + +func (ctx templateGo) genDefaultServiceTemplate() string { + return `package api + +import ( + "bufio" + "bytes" + "encoding/json" + "fmt" + "net/http" + "strings" +) + +// Result 统一响应结构 +type Result[T any] struct { + Id string + Code *uint16 + Data *T + Msg *string + Status uint8 +} + +func (r Result[T]) Error() string { + if r.Msg != nil { + return *r.Msg + } + if r.Code != nil { + return fmt.Sprintf("code=%d", *r.Code) + } + return "api: unknown error" +} + +// Void 用于无数据返回的方法 +type Void = struct{} + +// RequestInterceptor 请求前拦截器:可鉴权、加签、改 dto;返回 error 则直接失败 +type RequestInterceptor func(serviceName string, methodName string, dto any) error + +// ResponseInterceptor 响应后拦截器:可记录日志、埋点、解密;返回 error 则转为失败响应 +type ResponseInterceptor func(serviceName string, methodName string, result Result[any]) error + +// Client 内联 HTTP 客户端,不依赖外部 funclient 包 +type Client struct { + url string + client *http.Client + state map[string]string + requestInterceptors []RequestInterceptor + responseInterceptors []ResponseInterceptor +} + +// NewClient 创建客户端 +func NewClient(url string) (*Client, error) { + return &Client{ + url: strings.TrimRight(url, "/"), + client: &http.Client{}, + }, nil +} + +// SetHttpClient 替换底层 http.Client +func (c *Client) SetHttpClient(client *http.Client) { + c.client = client +} + +// AddRequestInterceptor 注册请求前拦截器 +func (c *Client) AddRequestInterceptor(i RequestInterceptor) { + c.requestInterceptors = append(c.requestInterceptors, i) +} + +// AddResponseInterceptor 注册响应后拦截器 +func (c *Client) AddResponseInterceptor(i ResponseInterceptor) { + c.responseInterceptors = append(c.responseInterceptors, i) +} + +// SetState 设置随每个请求携带的状态(如 token),服务端 Guard 可读取 +func (c *Client) SetState(state map[string]string) { + c.state = state +} + +// Request 发起普通调用 +func Request[T any](c *Client, serviceName string, methodName string, dto ...any) Result[T] { + payload := newPayload(serviceName, methodName, dto, c.state) + b, err := json.Marshal(payload) + if err != nil { + return Result[T]{Status: 2, Msg: ptr(err.Error())} + } + for _, i := range c.requestInterceptors { + var dtoVal any + if len(dto) > 0 { + dtoVal = dto[0] + } + if err := i(serviceName, methodName, dtoVal); err != nil { + return Result[T]{Status: 2, Msg: ptr(err.Error())} + } + } + req, err := http.NewRequest(http.MethodPost, c.url+"/cell", bytes.NewReader(b)) + if err != nil { + return Result[T]{Status: 2, Msg: ptr(err.Error())} + } + req.Header.Set("Content-Type", "application/json") + resp, err := c.client.Do(req) + if err != nil { + return Result[T]{Status: 2, Msg: ptr(err.Error())} + } + defer resp.Body.Close() + var out Result[T] + if err := json.NewDecoder(resp.Body).Decode(&out); err != nil { + return Result[T]{Status: 2, Msg: ptr(err.Error())} + } + var anyData *any + if out.Data != nil { + v := any(*out.Data) + anyData = &v + } + anyResult := Result[any]{Id: out.Id, Code: out.Code, Data: anyData, Msg: out.Msg, Status: out.Status} + for _, i := range c.responseInterceptors { + if err := i(serviceName, methodName, anyResult); err != nil { + return Result[T]{Status: 2, Msg: ptr(err.Error())} + } + } + return out +} + +// Stream 发起流式调用,通过 NDJSON 行逐个推送消息(Streamable HTTP) +func Stream[T any](c *Client, serviceName string, methodName string, dto ...any) (<-chan T, error) { + payload := newPayload(serviceName, methodName, dto, c.state) + b, err := json.Marshal(payload) + if err != nil { + return nil, err + } + for _, i := range c.requestInterceptors { + var dtoVal any + if len(dto) > 0 { + dtoVal = dto[0] + } + if err := i(serviceName, methodName, dtoVal); err != nil { + return nil, err + } + } + req, err := http.NewRequest(http.MethodPost, c.url+"/cell", bytes.NewReader(b)) + if err != nil { + return nil, err + } + req.Header.Set("Content-Type", "application/json") + resp, err := c.client.Do(req) + if err != nil { + return nil, err + } + if resp.StatusCode != http.StatusOK { + resp.Body.Close() + return nil, fmt.Errorf("api: unexpected status %d", resp.StatusCode) + } + anyResult := Result[any]{Status: 0} + for _, i := range c.responseInterceptors { + if err := i(serviceName, methodName, anyResult); err != nil { + resp.Body.Close() + return nil, err + } + } + ch := make(chan T) + go func() { + defer resp.Body.Close() + defer close(ch) + scanner := bufio.NewScanner(resp.Body) + for scanner.Scan() { + line := strings.TrimSpace(scanner.Text()) + if line == "" { + continue + } + var msg T + if err := json.Unmarshal([]byte(line), &msg); err != nil { + continue + } + ch <- msg + } + }() + return ch, nil +} + +func newPayload(serviceName string, methodName string, dto []any, state map[string]string) map[string]any { + payload := map[string]any{ + "serviceName": serviceName, + "methodName": methodName, + } + if len(dto) > 0 { + payload["data"] = dto[0] + } + if len(state) > 0 { + payload["state"] = state + } + return payload +} + +func ptr(s string) *string { return &s } + +type Api struct { +{{- range .GenServiceList}} + {{.ServiceName}} *{{.ServiceName}} +{{- end}} + *Client +} + +func CreateApi(url string) (Api, error) { + apiClient, err := NewClient(url) + return Api{ +{{- range .GenServiceList}} + {{.ServiceName}}: New{{.ServiceName}}(apiClient), +{{- end}} + Client: apiClient, + }, err +}` +} + +func (ctx templateGo) genServiceTemplate() string { + return `package api + +type {{.ServiceName}} struct { + *Client +} + +func New{{.ServiceName}}(client *Client) *{{.ServiceName}} { + return &{{.ServiceName}}{ + Client: client, + } +} + +{{- $serviceName := .ServiceName }} +{{- range .GenMethodTypeList}} +{{if .IsStream }}func (ctx *{{$serviceName}}) {{.MethodName}}({{.DtoText}}) (<-chan {{.GenericTypeText}}, error) { + return Stream[{{.GenericTypeText}}](ctx.Client, "{{$serviceName}}", "{{.MethodName}}"{{.ArgsText}}) +}{{else}}func (ctx *{{$serviceName}}) {{.MethodName}}({{.DtoText}}) {{.ReturnValueText}} { + return Request[{{.GenericTypeText}}](ctx.Client, "{{$serviceName}}", "{{.MethodName}}"{{.ArgsText}}) +}{{end}} +{{- end}}` +} + +func (ctx templateGo) genStructTemplate() string { + return `package api + +type {{.Name}} struct{ + {{- range .GenClassFieldType}} + {{.Name}} {{.Type}} {{.Tag}} + {{- end}} +}` +} + +func (ctx templateGo) genEnumTemplate() string { + return `package api + +type {{.Name}} uint8 + +{{$enumName := .Name}} +const ( +{{- range $index, $element := .Names}} + {{$element}}{{if eq $index 0}} {{$enumName}} = iota{{end}} +{{- end}} +) + +func ({{.Name}}) Values() []{{.Name}} { + return []{{.Name}}{ +{{- range $index, $element := .Names}} + {{$element}}, +{{- end}} + } +} + +{{if .DisplayNames}} +func ({{.Name}}) DisplayNames() []string { + return []string{ +{{- range $index, $element := .DisplayNames}} + "{{$element}}", +{{- end}} + } +} +{{end}}` +} diff --git a/template_ts.go b/template_ts.go new file mode 100644 index 0000000..364a302 --- /dev/null +++ b/template_ts.go @@ -0,0 +1,186 @@ +package fun + +type templateTs struct{} + +func (ctx templateTs) genClientTemplate() string { + return `export type result = { + id?: string; + code?: number; + data?: T; + msg?: string; + status: number; +}; + +export type RequestInterceptor = ( + serviceName: string, + methodName: string, + dto: any +) => Promise | void; + +export type ResponseInterceptor = ( + serviceName: string, + methodName: string, + result: result +) => Promise | void; + +export class Client { + private url: string; + private state: Record = {}; + private requestInterceptors: RequestInterceptor[] = []; + private responseInterceptors: ResponseInterceptor[] = []; + + constructor(url: string) { + this.url = url.replace(/\/+$/, ""); + } + + setState(state: Record) { + this.state = state; + } + + addRequestInterceptor(i: RequestInterceptor) { + this.requestInterceptors.push(i); + } + + addResponseInterceptor(i: ResponseInterceptor) { + this.responseInterceptors.push(i); + } + + async request(serviceName: string, methodName: string, dto?: any): Promise> { + for (const i of this.requestInterceptors) { + await i(serviceName, methodName, dto); + } + const res = await fetch(this.url + "/cell", { + method: "POST", + headers: { "Content-Type": "application/json" }, + body: JSON.stringify({ serviceName, methodName, data: dto, ...(Object.keys(this.state).length ? { state: this.state } : {}) }), + }); + const out = (await res.json()) as result; + const anyResult: result = { + id: out.id, + code: out.code, + data: out.data, + msg: out.msg, + status: out.status, + }; + for (const i of this.responseInterceptors) { + await i(serviceName, methodName, anyResult); + } + return out; + } + + async stream( + serviceName: string, + methodName: string, + dto: any | undefined, + onMessage: (data: T) => void + ): Promise { + for (const i of this.requestInterceptors) { + await i(serviceName, methodName, dto); + } + const res = await fetch(this.url + "/cell", { + method: "POST", + headers: { "Content-Type": "application/json" }, + body: JSON.stringify({ serviceName, methodName, data: dto, ...(Object.keys(this.state).length ? { state: this.state } : {}) }), + }); + if (!res.ok) return; + const anyResult: result = { status: 0 }; + for (const i of this.responseInterceptors) { + await i(serviceName, methodName, anyResult); + } + if (!res.body) return; + const reader = res.body.getReader(); + const decoder = new TextDecoder(); + let buffer = ""; + for (;;) { + const { done, value } = await reader.read(); + if (done) break; + buffer += decoder.decode(value, { stream: true }); + const lines = buffer.split("\n"); + buffer = lines.pop() ?? ""; + for (const line of lines) { + const payload = line.trim(); + if (!payload) continue; + const data = JSON.parse(payload) as T; + onMessage(data); + } + } + } +}` +} + +func (ctx templateTs) genDefaultServiceTemplate() string { + return `import { Client, type result } from "./client"; +{{- range .GenServiceList}} +import {{.ServiceName}} from "./{{.ServiceName}}"; +{{- end}} + +export class defaultApi extends Client { + constructor(url: string) { + super(url); + } + {{- range .GenServiceList}} + public {{.ServiceName}}: {{.ServiceName}} = new {{.ServiceName}}(this); + {{- end}} +} + +export default class api { + static create(url: string): defaultApi { + return new defaultApi(url); + } +}` +} + +func (ctx templateTs) genServiceTemplate() string { + return `import { Client, type result } from "./client" +{{- range .GenImport}} +import type {{.Name}} from "./{{.Name}}"; +{{- end}} + +export default class {{.ServiceName}} { + private client: Client; + constructor(client: Client) { + this.client = client; + } + {{- $serviceName := .ServiceName }} + {{- range .GenMethodTypeList}} + {{if .IsStream }}async {{.MethodName}}({{.DtoText}}{{if .DtoText}},{{end}}onMessage: (data: {{.GenericTypeText}}) => void): Promise { + return await this.client.stream<{{.GenericTypeText}}>("{{$serviceName}}", "{{.MethodName}}", {{if .DtoText}}dto{{else}}undefined{{end}}, onMessage) + }{{else}}async {{.MethodName}}({{.DtoText}}): Promise<{{.ReturnValueText}}> { + return await this.client.request<{{.GenericTypeText}}>("{{$serviceName}}", "{{.MethodName}}"{{.ArgsText}}) + }{{end}} + {{- end}} +}` +} + +func (ctx templateTs) genStructTemplate() string { + return `{{- range .GenImport}}import type {{.Name}} from "./{{.Name}}";{{"\n"}}{{- end}}export default interface {{.Name}} { + {{- range .GenClassFieldType}} + {{.Name}}:{{.Type}} + {{- end}} +}` +} + +func (ctx templateTs) genEnumTemplate() string { + return `enum {{.Name}} { +{{- range $index, $element := .Names}} + {{$element}}, +{{- end}} +}{{ $enumName := .Name }} +function values(): {{.Name}}[] { + return [ +{{- range $index, $element := .Names}} + {{$enumName}}.{{$element}}, +{{- end}} + ] +} +{{if .DisplayNames}} +function displayNames(): string[] { + return [ +{{- range $index, $element := .DisplayNames}} + "{{$element}}", +{{- end}} + ] +} +{{end}} +export default {{.Name}}` +} diff --git a/text.go b/text.go new file mode 100644 index 0000000..1d7bc6a --- /dev/null +++ b/text.go @@ -0,0 +1,25 @@ +// 文本处理辅助方法 + +package fun + +import "unicode" + +// 首字母转大写 +func firstLetterToUpper(s string) string { + if len(s) == 0 { + return s + } + runes := []rune(s) + runes[0] = unicode.ToUpper(runes[0]) + return string(runes) +} + +// 首字母转小写 +func firstLetterToLower(s string) string { + if len(s) == 0 { + return s + } + runes := []rune(s) + runes[0] = unicode.ToLower(runes[0]) + return string(runes) +} diff --git a/utils.go b/utils.go new file mode 100644 index 0000000..9830771 --- /dev/null +++ b/utils.go @@ -0,0 +1,60 @@ +package fun + +import ( + "net" + "net/http" + "strings" +) + +// getIP 获取客户端真实 IP +// 优先级:X-Forwarded-For > X-Real-IP > RemoteAddr +func getIP(r *http.Request) string { + // 1. 优先获取真实 IP(多层代理时取最后一个非空段) + if ip := lastNonEmpty(r.Header.Get("X-Forwarded-For")); ip != "" { + return toLoopback(ip) + } + + // 2. X-Real-IP(通常由 Nginx 设置) + if ip := strings.TrimSpace(r.Header.Get("X-Real-IP")); ip != "" { + return toLoopback(ip) + } + + // 3. 最终回退到 RemoteAddr(兼容带端口、IPv6 方括号、无端口) + if ip := hostOf(r.RemoteAddr); ip != "" { + return toLoopback(ip) + } + + return "127.0.0.1" +} + +// lastNonEmpty 取 X-Forwarded-For 中最后一个非空段 +// X-Forwarded-For: client, proxy1, proxy2 +func lastNonEmpty(xff string) string { + parts := strings.Split(xff, ",") + for i := len(parts) - 1; i >= 0; i-- { + if ip := strings.TrimSpace(parts[i]); ip != "" { + return ip + } + } + return "" +} + +// hostOf 从 RemoteAddr 中提取 IP 部分 +// "203.0.113.9:4567" → "203.0.113.9","[::1]:4567" → "::1", +// "198.51.100.88"(无端口)→ 原样返回 +func hostOf(remoteAddr string) string { + raw := strings.TrimSpace(remoteAddr) + host, _, err := net.SplitHostPort(raw) + if err == nil && host != "" { + return host + } + return raw +} + +// toLoopback 回环地址统一返回 127.0.0.1,其余原样返回 +func toLoopback(ip string) string { + if parsed := net.ParseIP(ip); parsed != nil && parsed.IsLoopback() { + return "127.0.0.1" + } + return ip +}