diff --git a/benchmarks_test.go b/benchmarks_test.go index 41f6d72..e6d3ec2 100644 --- a/benchmarks_test.go +++ b/benchmarks_test.go @@ -23,6 +23,7 @@ func setupGoRPC(b *testing.B) *server.ServiceGoRPCClient { if err != nil { b.Fatal(err) } + addr := l.Addr().String() if err := l.Close(); err != nil { b.Fatal(err) @@ -32,6 +33,7 @@ func setupGoRPC(b *testing.B) *server.ServiceGoRPCClient { if err := s.Start(); err != nil { b.Fatal(err) } + b.Cleanup(s.Stop) c := server.NewServiceGoRPCClient(addr, nil) @@ -43,8 +45,10 @@ func setupGoRPC(b *testing.B) *server.ServiceGoRPCClient { func setupGoTSRPC(b *testing.B) *server.HTTPServiceGoTSRPCClient { b.Helper() + s := httptest.NewServer(server.NewDefaultServiceGoTSRPCProxy(&server.Handler{})) b.Cleanup(s.Close) + return server.NewDefaultServiceGoTSRPCClient(s.URL) } @@ -72,6 +76,7 @@ func Benchmark_Empty(b *testing.B) { func Benchmark_String(b *testing.B) { b.ReportAllocs() + v := "hello world" b.Run("GoRPC", func(b *testing.B) { @@ -95,6 +100,7 @@ func Benchmark_String(b *testing.B) { func Benchmark_SimpleStruct(b *testing.B) { b.ReportAllocs() + v := common.Simple{ Bool: true, Int: 42, @@ -124,6 +130,7 @@ func Benchmark_SimpleStruct(b *testing.B) { func Benchmark_NestedStruct(b *testing.B) { b.ReportAllocs() + v := common.Nested{ Name: "parent", Child: common.Simple{ @@ -156,6 +163,7 @@ func Benchmark_NestedStruct(b *testing.B) { func Benchmark_StructWithCollections(b *testing.B) { b.ReportAllocs() + v := server.WithCollections{ Strings: []string{"a", "b", "c"}, Int64s: []int64{1, 2, 3}, diff --git a/bufferedclient.go b/bufferedclient.go index 6bb9e0d..037604c 100644 --- a/bufferedclient.go +++ b/bufferedclient.go @@ -30,6 +30,7 @@ func (c *bufferedClient) SetTransportHttpClient(client *http.Client) { //nolint: // Call calls a method on the remote service func (c *bufferedClient) Call(ctx context.Context, url string, endpoint string, method string, args []any, reply []any) error { var errorIndices []int + for i, v := range reply { if isErrorPtr(v) { errorIndices = append(errorIndices, i) @@ -40,9 +41,11 @@ func (c *bufferedClient) Call(ctx context.Context, url string, endpoint string, if len(args) > 0 { b = getBuffer() defer putBuffer(b) + enc := c.handle.getEncoder(b) err := enc.Encode(args) c.handle.putEncoder(enc) + if err != nil { return NewClientError(errors.Wrap(err, "failed to encode arguments")) } @@ -56,6 +59,7 @@ func (c *bufferedClient) Call(ctx context.Context, url string, endpoint string, if c.headers != nil { headers = c.headers.Clone() } + request, errRequest := newRequest(ctx, postURL, c.handle.contentType, b, headers) if errRequest != nil { return NewClientError(errors.Wrap(errRequest, "failed to create request")) @@ -69,6 +73,7 @@ func (c *bufferedClient) Call(ctx context.Context, url string, endpoint string, buf := getBuffer() defer putBuffer(buf) + if _, err := io.Copy(buf, resp.Body); err != nil { return NewClientError(errors.Wrap(err, "failed to read response body")) } @@ -95,6 +100,7 @@ func (c *bufferedClient) Call(ctx context.Context, url string, endpoint string, dec := clientHandle.getDecoder(buf) err := dec.Decode(wrappedReply) clientHandle.putDecoder(dec) + if err != nil { return NewClientError(errors.Wrap(err, "failed to decode response")) } diff --git a/client.go b/client.go index 6d6fcf6..c67651a 100644 --- a/client.go +++ b/client.go @@ -33,13 +33,16 @@ func newRequest(ctx context.Context, url string, contentType string, buffer *byt if buffer == nil { buffer = &bytes.Buffer{} } + request, errRequest := http.NewRequestWithContext(ctx, http.MethodPost, url, buffer) if errRequest != nil { return nil, errors.Wrap(errRequest, "could not create a request") } + if len(headers) > 0 { request.Header = headers } + request.Header.Set("Content-Type", contentType) request.Header.Set("Accept", contentType) request.Header.Set(HeaderServiceToService, "true") diff --git a/clienterror.go b/clienterror.go index 4610468..39e1b45 100644 --- a/clienterror.go +++ b/clienterror.go @@ -15,5 +15,6 @@ func (e *ClientError) Unwrap() error { if e != nil && e.error != nil { return e.error } + return nil } diff --git a/cmd/gotsrpc/gotsrpc.go b/cmd/gotsrpc/gotsrpc.go index 10084e2..fc2c950 100644 --- a/cmd/gotsrpc/gotsrpc.go +++ b/cmd/gotsrpc/gotsrpc.go @@ -43,6 +43,7 @@ func main() { flagDebug := flag.Bool("debug", false, "debug") flag.Usage = usage + flag.Parse() ctx := context.Background() @@ -58,10 +59,12 @@ func main() { if value, err := strconv.ParseInt(buildTimestamp, 10, 64); err == nil { buildTime = time.Unix(value, 0).String() } + fmt.Printf("Version: %s\nCommit: %s\nBuildTime: %s\n", version, commitHash, buildTime) } else { fmt.Println(version) } + os.Exit(0) case len(args) != 1: usage() @@ -71,14 +74,18 @@ func main() { codegen.Trace = *flagDebug } - var goRoot string - var goPath string + var ( + goRoot string + goPath string + ) + if out, err := exec.CommandContext(ctx, "go", "env", "GOROOT").Output(); err != nil { fmt.Println("failed to retrieve GOROOT", err.Error()) os.Exit(1) } else { goRoot = string(bytes.TrimSpace(out)) } + if out, err := exec.CommandContext(ctx, "go", "env", "GOPATH").Output(); err != nil { fmt.Println("failed to retrieve GOPATH", err.Error()) os.Exit(1) @@ -89,6 +96,7 @@ func main() { conf, err := config.LoadConfigFile(args[0]) if err != nil { _, _ = fmt.Fprintln(os.Stderr, "config load error, could not load config from", args[0], ":", err) + os.Exit(2) } diff --git a/config/config.go b/config/config.go index 908aa67..f04c561 100644 --- a/config/config.go +++ b/config/config.go @@ -33,6 +33,7 @@ func (t *Target) IsGoRPC(service string) bool { return true } } + return false } @@ -40,11 +41,13 @@ func (t *Target) IsTSRPC(service string) bool { if len(t.TSRPC) == 0 { return true } + for _, value := range t.TSRPC { if value == service { return true } } + return false } @@ -86,6 +89,7 @@ func LoadConfigFile(file string) (conf *Config, err error) { if readErr != nil { return nil, errors.New("could not read config file: " + readErr.Error()) } + conf, err = loadConfig(yamlBytes) if err != nil { return nil, err @@ -96,6 +100,7 @@ func LoadConfigFile(file string) (conf *Config, err error) { if err != nil { return nil, err } + conf.Module.Path = absPath if data, err := os.ReadFile(path.Join(absPath, "go.mod")); err != nil && !os.IsNotExist(err) { @@ -105,21 +110,26 @@ func LoadConfigFile(file string) (conf *Config, err error) { if err != nil { return nil, err } + conf.Module.ModFile = modFile } } + return conf, nil } func loadConfig(yamlBytes []byte) (conf *Config, err error) { conf = &Config{} + yamlErr := yaml.Unmarshal(yamlBytes, conf) if yamlErr != nil { err = errors.New("could not parse yaml: " + yamlErr.Error()) return } + for goPackage, mapping := range conf.Mappings { mapping.GoPackage = goPackage } + return } diff --git a/config/config_test.go b/config/config_test.go index fad3da1..ef09ce7 100644 --- a/config/config_test.go +++ b/config/config_test.go @@ -33,7 +33,9 @@ func TestLoadConfig(t *testing.T) { if err != nil { t.Fatal(err) } + goPackage := "foo/bar" + foo, ok := c.Mappings[goPackage] if !ok { t.Fatal("foo/bar not found") @@ -42,6 +44,7 @@ func TestLoadConfig(t *testing.T) { if foo.GoPackage != goPackage { t.Fatal("wrong go package value") } + if foo.Out != "path/to/ts" || foo.TypeScriptModule != "Sample.Module" { t.Fatal("unexpected data", foo) } @@ -52,18 +55,23 @@ func TestLoadConfig(t *testing.T) { if !ok { t.Fatal("demo target not found") } + if demoTarget.Out != "/tmp/my-service.ts" { t.Fatal("demo target out is wrong") } + if demoTarget.Package != "github.com/foomo/gotsrpc/v2/demo" { t.Fatal("wrong target package") } + if demoTarget.TypeScriptModule != "My.Service" { t.Fatal("wromg ts module") } + if len(demoTarget.Services) != 1 { t.Fatal("wrong number of services") } + if demoTarget.Services["/service/demo"] != "Service" { t.Fatal("first service is wrong") } diff --git a/context.go b/context.go index 7543820..487a613 100644 --- a/context.go +++ b/context.go @@ -18,6 +18,7 @@ func GetStatsForRequest(r *http.Request) (*CallStats, bool) { if value, ok := r.Context().Value(contextStatsKey).(*CallStats); ok && value != nil { return value, true } + return &CallStats{}, false } diff --git a/error.go b/error.go index f5ad461..5a33514 100644 --- a/error.go +++ b/error.go @@ -40,6 +40,7 @@ func NewError(err error) *Error { // retrieve error details errType := reflect.TypeOf(err) + errElem := errType if errType.Kind() == reflect.Ptr { errElem = errType.Elem() @@ -60,6 +61,7 @@ func NewError(err error) *Error { for i, e := range errs { inst.ErrCauses[i] = NewError(e) } + return inst } } @@ -78,6 +80,7 @@ func (e *Error) As(err interface{}) bool { if e == nil || err == nil { return false } + if reflect.TypeOf(err).Elem().String() == e.Type { if decodeErr := mapstructure.Decode(e.Data, &err); decodeErr != nil { fmt.Printf("ERROR: failed to decode error data\n%+v", decodeErr) @@ -86,6 +89,7 @@ func (e *Error) As(err interface{}) bool { return true } } + return false } @@ -94,6 +98,7 @@ func (e *Error) Cause() error { if e.ErrCause != nil { return e.ErrCause } + return e } @@ -107,6 +112,7 @@ func (e *Error) Format(s fmt.State, verb rune) { _, _ = fmt.Fprintf(s, "Data: %v\n", e.Data) } } + fallthrough case 's', 'q': _, _ = io.WriteString(s, e.Error()) @@ -118,18 +124,22 @@ func (e *Error) Unwrap() []error { if e == nil { return nil } + var errs []error if e.ErrCause != nil { errs = append(errs, e.ErrCause) } + for _, c := range e.ErrCauses { if c != nil { errs = append(errs, c) } } + if len(errs) == 0 { return nil } + return errs } @@ -140,6 +150,7 @@ func (e *Error) Is(err error) bool { } errType := reflect.TypeOf(err) + errElem := errType if errType.Kind() == reflect.Ptr { errElem = errType.Elem() @@ -160,5 +171,6 @@ func (e *Error) Error() string { if e.ErrCause != nil { msg += ": " + e.ErrCause.Error() } + return msg } diff --git a/example/basic/main.go b/example/basic/main.go index e9e608a..b6f83ae 100644 --- a/example/basic/main.go +++ b/example/basic/main.go @@ -27,6 +27,7 @@ func main() { go func() { time.Sleep(time.Second) + _ = exec.CommandContext(ctx, "open", "http://127.0.0.1:3000").Run() }() diff --git a/example/basic/main_test.go b/example/basic/main_test.go index aa6c1ee..1ab6097 100644 --- a/example/basic/main_test.go +++ b/example/basic/main_test.go @@ -12,6 +12,7 @@ import ( func TestContextCanceled(t *testing.T) { ctx, cancel := context.WithCancel(t.Context()) + go func() { time.Sleep(time.Second) cancel() diff --git a/example/monitor/main.go b/example/monitor/main.go index 2ea3724..906e4e9 100644 --- a/example/monitor/main.go +++ b/example/monitor/main.go @@ -42,6 +42,7 @@ func main() { go func() { time.Sleep(time.Second) + _ = exec.CommandContext(ctx, "open", "http://127.0.0.1:3000").Run() call(ctx) }() diff --git a/gotsrpc.go b/gotsrpc.go index 150b97c..25ea600 100644 --- a/gotsrpc.go +++ b/gotsrpc.go @@ -42,13 +42,16 @@ func LoadArgs(args interface{}, callStats *CallStats, r *http.Request) error { dec := ch.getDecoder(r.Body) errDecode := dec.Decode(args) ch.putDecoder(dec) + if errDecode != nil { return errors.Wrap(errDecode, "could not decode arguments") } + if callStats != nil { callStats.Unmarshalling = time.Since(start) callStats.RequestSize = int(r.ContentLength) } + return nil } @@ -56,18 +59,21 @@ func loadArgs(args interface{}, jsonBytes []byte) error { if err := json.Unmarshal(jsonBytes, &args); err != nil { return err } + return nil } // Reply although this is a public method - do not call it, it will be called by generated code func Reply(response []interface{}, stats *CallStats, r *http.Request, w http.ResponseWriter) error { var errorIndices []int + for i, v := range response { if er, ok := v.(*errorReply); ok { errorIndices = append(errorIndices, i) response[i] = er.err } } + var serializationStart time.Time if stats != nil { serializationStart = time.Now() @@ -89,6 +95,7 @@ func Reply(response []interface{}, stats *CallStats, r *http.Request, w http.Res enc := ch.getEncoder(buf) err := enc.Encode(response) ch.putEncoder(enc) + if err != nil { return errors.Wrap(err, "could not encode data to accepted format") } @@ -102,11 +109,13 @@ func Reply(response []interface{}, stats *CallStats, r *http.Request, w http.Res if stats != nil { stats.ResponseSize = buf.Len() stats.Marshalling = time.Since(serializationStart) + for _, i := range errorIndices { if v, ok := response[i].(error); ok && v != nil { if !reflect.ValueOf(v).IsZero() { stats.ErrorCode = 1 stats.ErrorType = fmt.Sprintf("%T", v) + stats.ErrorMessage = v.Error() if v, ok := v.(interface { ErrorCode() int @@ -117,5 +126,6 @@ func Reply(response []interface{}, stats *CallStats, r *http.Request, w http.Res } } } + return nil } diff --git a/gotsrpc_test.go b/gotsrpc_test.go index 7622fdd..9dcb419 100644 --- a/gotsrpc_test.go +++ b/gotsrpc_test.go @@ -9,16 +9,20 @@ func TestLoadArgs(t *testing.T) { foo := "" bar := []string{} args := []interface{}{&foo, &bar} + errLoad := loadArgs(&args, jsonBytes) if errLoad != nil { t.Fatal(errLoad) } + if foo != "a" { t.Fatal("foo should have been a") } + if len(bar) != 3 { t.Fatal("bar len wrong", len(bar), "!=", len(bar)) } + if bar[1] != "b" { t.Fatal("bar[1] (", bar[1], ") != b") } @@ -26,25 +30,32 @@ func TestLoadArgs(t *testing.T) { func TestLoadInterfaceArgs(t *testing.T) { jsonBytes := []byte(`["a", ["a", "b", "c"], 1.3]`) + var ( foo interface{} bar []interface{} floaty interface{} ) + args := []interface{}{&foo, &bar, &floaty} + errLoad := loadArgs(&args, jsonBytes) if errLoad != nil { t.Fatal(errLoad) } + if foo != "a" { t.Fatal("foo should have been a") } + if len(bar) != 3 { t.Fatal("bar len wrong", len(bar), "!=", len(bar)) } + if bar[1] != "b" { t.Fatal("bar[1] (", bar[1], ") != b") } + if floaty != 1.3 { t.Fatal("floaty mismatch", floaty) } diff --git a/instrumentation.go b/instrumentation.go index 81eced8..83ca545 100644 --- a/instrumentation.go +++ b/instrumentation.go @@ -7,6 +7,7 @@ func InstrumentedService(middleware http.HandlerFunc, handleStats GoRPCCallStats return func(w http.ResponseWriter, r *http.Request) { *r = *RequestWithStatsContext(r) middleware(w, r) + if stats, ok := GetStatsForRequest(r); ok { handleStats(stats) } diff --git a/instrumentation_test.go b/instrumentation_test.go index 7645fb5..7828fe7 100644 --- a/instrumentation_test.go +++ b/instrumentation_test.go @@ -24,6 +24,7 @@ func TestInstrumentedService(t *testing.T) { assert.Equal(t, "package", s.Package) assert.Equal(t, "service", s.Service) assert.NotNil(t, s) + count++ }) diff --git a/internal/build/build.go b/internal/build/build.go index 7fc3b72..2103923 100644 --- a/internal/build/build.go +++ b/internal/build/build.go @@ -28,9 +28,11 @@ func Build(conf *config.Config, goPath, goRoot string) { //nolint:maintidx for name := range conf.Targets { names = append(names, name) } + sort.Strings(names) missingTypes := map[string]bool{} + for _, mapping := range conf.Mappings { for _, include := range mapping.Structs { missingTypes[include] = true @@ -38,6 +40,7 @@ func Build(conf *config.Config, goPath, goRoot string) { //nolint:maintidx } missingConstants := map[string]bool{} + for _, mapping := range conf.Mappings { for _, include := range mapping.Scalars { missingConstants[include] = true @@ -71,8 +74,10 @@ func Build(conf *config.Config, goPath, goRoot string) { //nolint:maintidx workDirectory, err := os.Getwd() if err != nil { _, _ = fmt.Fprintln(os.Stderr, err) + os.Exit(1) } + vendorDirectory := path.Join(workDirectory, "vendor") goPaths := []string{goPath, goRoot} @@ -84,11 +89,13 @@ func Build(conf *config.Config, goPath, goRoot string) { //nolint:maintidx pkgName, services, structs, scalars, constantTypes, err := parser.Read(goPaths, conf.Module, packageName, target.Services, missingTypes, missingConstants) if err != nil { _, _ = fmt.Fprintln(os.Stderr, "\t an error occurred while trying to understand your code: ", err) + os.Exit(2) } // collect all union structs unions := map[string][]string{} + for _, s := range structs { if len(s.Fields) == 0 && len(s.UnionFields) > 0 { unions[s.Package] = append(unions[s.Package], s.Name) @@ -99,6 +106,7 @@ func Build(conf *config.Config, goPath, goRoot string) { //nolint:maintidx ts, err := codegen.RenderTypeScriptServices(services, conf.Mappings, scalars, structs, target) if err != nil { _, _ = fmt.Fprintln(os.Stderr, " could not generate ts code", err) + os.Exit(3) } @@ -111,12 +119,14 @@ func Build(conf *config.Config, goPath, goRoot string) { //nolint:maintidx updateErr := updateCode(target.Out, codegen.GetTSHeaderComment()+ts) if updateErr != nil { _, _ = fmt.Fprintln(os.Stderr, " could not write service file", target.Out, updateErr) + os.Exit(3) } err = codegen.RenderTypescriptStructsToPackages(structs, conf.Mappings, constantTypes, scalars, mappedTypeScript) if err != nil { _, _ = fmt.Fprintln(os.Stderr, "struct gen err for target", name, err) + os.Exit(4) } } @@ -140,8 +150,10 @@ func Build(conf *config.Config, goPath, goRoot string) { //nolint:maintidx writeErr := os.WriteFile(filename, []byte(code), 0644) //nolint:gosec if writeErr != nil { _, _ = fmt.Fprintln(os.Stderr, " could not write go source to file", writeErr) + os.Exit(5) } + _, _ = fmt.Fprintln(os.Stderr, "wrote code for debugging into file", filename) os.Exit(5) @@ -150,23 +162,30 @@ func Build(conf *config.Config, goPath, goRoot string) { //nolint:maintidx writeErr := os.WriteFile(filename, codeBytes, 0644) //nolint:gosec if writeErr != nil { _, _ = fmt.Fprintln(os.Stderr, " could not write go source to file", writeErr) + os.Exit(5) } } + if len(target.TSRPC) > 0 { goTSRPCProxiesCode, goerr := codegen.RenderGoTSRPCProxies(services, packageName, pkgName, target, unions) if goerr != nil { _, _ = fmt.Fprintln(os.Stderr, " could not generate go ts rpc proxies code in target", name, goerr) + os.Exit(4) } + formatAndWrite(goTSRPCProxiesCode, goTSRPCProxiesFilename) } + if len(target.TSRPC) > 0 && !target.SkipTSRPCClient { goTSRPCClientsCode, goerr := codegen.RenderGoTSRPCClients(services, packageName, pkgName, target) if goerr != nil { _, _ = fmt.Fprintln(os.Stderr, " could not generate go ts rpc clients code in target", name, goerr) + os.Exit(4) } + formatAndWrite(goTSRPCClientsCode, goTSRPCClientsFilename) } @@ -174,15 +193,19 @@ func Build(conf *config.Config, goPath, goRoot string) { //nolint:maintidx goRPCProxiesCode, goerr := codegen.RenderGoRPCProxies(services, packageName, pkgName, target) if goerr != nil { _, _ = fmt.Fprintln(os.Stderr, " could not generate go rpc proxies code in target", name, goerr) + os.Exit(4) } + formatAndWrite(goRPCProxiesCode, goRPCProxiesFilename) goRPCClientsCode, goerr := codegen.RenderGoRPCClients(services, packageName, pkgName, target) if goerr != nil { _, _ = fmt.Fprintln(os.Stderr, " could not generate go rpc clients code in target", name, goerr) + os.Exit(4) } + formatAndWrite(goRPCClientsCode, goRPCClientsFilename) } } @@ -191,6 +214,7 @@ func Build(conf *config.Config, goPath, goRoot string) { //nolint:maintidx mapping, ok := conf.Mappings[goPackage] if !ok { _, _ = fmt.Fprintln(os.Stderr, "reverse mapping error in struct generation for package", goPackage) + os.Exit(6) } @@ -206,10 +230,12 @@ func Build(conf *config.Config, goPath, goRoot string) { //nolint:maintidx // sort and keep enums on top slices.SortFunc(structNames, func(e1 string, e2 string) int { es1, ok1 := mappedStructsMap[e1] + es2, ok2 := mappedStructsMap[e2] if !ok1 || !ok2 { return strings.Compare(e1, e2) } + es1E := strings.Contains(es1.String(), "export enum ") es2E := strings.Contains(es2.String(), "export enum ") @@ -222,6 +248,7 @@ func Build(conf *config.Config, goPath, goRoot string) { //nolint:maintidx return strings.Compare(e1, e2) } }) + for _, structName := range structNames { structCode, ok := mappedStructsMap[structName] if ok { @@ -256,7 +283,9 @@ func relativeFilePath(a, b string) (r string, e error) { if e != nil { return } + r = strings.TrimSuffix(r, ".ts") + return } @@ -265,7 +294,9 @@ func commonJSImports(conf *config.Config, c *codegen.Code, tsFilename string, co for packageName := range conf.Mappings { packageNames = append(packageNames, packageName) } + sort.Strings(packageNames) + for _, packageName := range packageNames { importMapping := conf.Mappings[packageName] @@ -278,6 +309,7 @@ func commonJSImports(conf *config.Config, c *codegen.Code, tsFilename string, co fmt.Println("can not derive a relative path between", tsFilename, "and", importMapping.Out, relativeErr) os.Exit(1) } + c.L("import * as " + importMapping.TypeScriptModule + " from './" + relativePath + "'; // " + tsFilename + " to " + importMapping.Out) } } @@ -298,18 +330,23 @@ func updateCode(file string, code string) error { if len(home) == 0 { return errors.New("could not resolve home dir") } + file = path.Join(home, file[1:]) } } + errMkdirAll := os.MkdirAll(path.Dir(file), 0755) //nolint:gosec if errMkdirAll != nil { return errMkdirAll } + oldCode, _ := os.ReadFile(file) //nolint:gosec if string(oldCode) != code { fmt.Println(" writing file", file) return os.WriteFile(file, []byte(code), 0600) //nolint:gosec } + fmt.Println(" update file not necessary - unchanged", file) + return nil } diff --git a/internal/codegen/code.go b/internal/codegen/code.go index 3412072..a2c2563 100644 --- a/internal/codegen/code.go +++ b/internal/codegen/code.go @@ -23,12 +23,14 @@ func (c *Code) Ind(inc int) *Code { if c.indent < 0 { c.indent = 0 } + return c } func (c *Code) NL() *Code { c.lines = append(c.lines, strings.Repeat(c.tab, c.indent)+c.line) c.line = "" + return c } @@ -47,5 +49,6 @@ func (c *Code) String() string { c.lines = append(c.lines, c.line) c.line = "" } + return strings.Join(c.lines, "\n") } diff --git a/internal/codegen/gocode.go b/internal/codegen/gocode.go index 26d9116..aaa0df8 100644 --- a/internal/codegen/gocode.go +++ b/internal/codegen/gocode.go @@ -28,6 +28,7 @@ func valueGoType(v *model.Value, aliases map[string]string, packageName string) if v.IsPtr { t = "*" } + switch { case v.Array != nil: t += "[]" + valueGoType(v.Array.Value, aliases, packageName) @@ -37,6 +38,7 @@ func valueGoType(v *model.Value, aliases map[string]string, packageName string) if packageName != v.StructType.Package && aliases[v.StructType.Package] != "" { t += aliases[v.StructType.Package] + "." } + t += v.StructType.Name case v.Map != nil: t += `map[` + valueGoType(v.Map.Key, aliases, packageName) + `]` + valueGoType(v.Map.Value, aliases, packageName) @@ -44,6 +46,7 @@ func valueGoType(v *model.Value, aliases map[string]string, packageName string) if packageName != v.Scalar.Package && aliases[v.Scalar.Package] != "" { t += aliases[v.Scalar.Package] + "." } + t += v.Scalar.Name case v.IsInterface: t += "any" @@ -64,6 +67,7 @@ func ucfirst(str string) string { func strfirst(str string, strfunc func(string) string) string { res := "" + for i, char := range str { if i == 0 { res += strfunc(string(char)) @@ -71,23 +75,27 @@ func strfirst(str string, strfunc func(string) string) string { res += string(char) } } + return res } func extractImport(packageName string, fullPackageName string, aliases map[string]string) { r := strings.NewReplacer(".", "_", "/", "_", "-", "_") + if packageName != fullPackageName { if _, ok := aliases[packageName]; !ok { packageParts := strings.Split(packageName, "/") beautifulAlias := packageParts[len(packageParts)-1] uglyAlias := r.Replace(packageName) alias := uglyAlias + for _, otherAlias := range aliases { if otherAlias == beautifulAlias { alias = uglyAlias break } } + aliases[packageName] = alias } } @@ -119,10 +127,12 @@ func renderTSRPCServiceProxies(services model.ServiceList, fullPackageName strin "net/http": "http", "github.com/foomo/gotsrpc/v2": "gotsrpc", } + for _, service := range services { if !config.IsTSRPC(service.Name) { continue } + for _, m := range service.Methods { extractImports(m.Args, fullPackageName, aliases) } @@ -149,9 +159,11 @@ func renderTSRPCServiceProxies(services model.ServiceList, fullPackageName strin proxyName := service.Name + "GoTSRPCProxy" g.L("const (") + for _, method := range service.Methods { g.L(proxyName + method.Name + " = \"" + method.Name + "\"") } + g.L(")") g.L(` @@ -199,20 +211,25 @@ func renderTSRPCServiceProxies(services model.ServiceList, fullPackageName strin for _, method := range service.Methods { g.L("case " + proxyName + method.Name + ":") g.Ind(1) + var ( callArgs []string isContextRequest bool isSessionRequest bool ) + g.L("var (") g.Ind(1) g.L("args []any") g.L("rets []any") g.Ind(-1) g.L(")") + if len(method.Args) > 0 { - var args []string - var argsDecls []string + var ( + args []string + argsDecls []string + ) skipArgI := 0 @@ -228,11 +245,14 @@ func renderTSRPCServiceProxies(services model.ServiceList, fullPackageName strin callArgs = append(callArgs, argName) skipArgI++ } + if len(args) > 0 { g.L("var (") + for _, argDecl := range argsDecls { g.L(argDecl) } + g.L(")") g.L("args = []any{" + strings.Join(args, ", ") + "}") g.L("if err := gotsrpc.LoadArgs(&args, callStats, r); err != nil {") @@ -243,7 +263,9 @@ func renderTSRPCServiceProxies(services model.ServiceList, fullPackageName strin g.L("}") } } + var returnValueNames []string + for retI, retField := range method.Return { retArgName := retField.Name if len(retArgName) == 0 { @@ -252,8 +274,10 @@ func renderTSRPCServiceProxies(services model.ServiceList, fullPackageName strin retArgName += "_" + fmt.Sprint(retI) } } + returnValueNames = append(returnValueNames, lcfirst(method.Name)+ucfirst(retArgName)) } + g.L("var executionStart time.Time") g.L("if callStats != nil {") g.Ind(1) @@ -263,13 +287,16 @@ func renderTSRPCServiceProxies(services model.ServiceList, fullPackageName strin if isSessionRequest { g.L("rw := gotsrpc.ResponseWriter{ResponseWriter: w}") + callArgs = append([]string{"&rw", "r"}, callArgs...) } else if isContextRequest { callArgs = append([]string{"r.Context()"}, callArgs...) } + if len(returnValueNames) > 0 { g.App(strings.Join(returnValueNames, ", ") + " := ") } + g.App("p.service." + method.Name + "(" + strings.Join(callArgs, ", ") + ")") g.NL() g.L("if callStats != nil {") @@ -277,17 +304,20 @@ func renderTSRPCServiceProxies(services model.ServiceList, fullPackageName strin g.L("callStats.Execution = time.Since(executionStart)") g.Ind(-1) g.L("}") + if isSessionRequest { g.L("if rw.Status() == http.StatusOK {").Ind(1) } // Wrap all error return values with ErrorReply retsValues := make([]string, len(returnValueNames)) copy(retsValues, returnValueNames) + for i, ret := range method.Return { if ret.Value.GoScalarType == "error" { retsValues[i] = "gotsrpc.ErrorReply(" + retsValues[i] + ")" } } + g.L("rets = []any{" + strings.Join(retsValues, ", ") + "}") g.L("if err := gotsrpc.Reply(rets, callStats, r, w); err != nil {") g.Ind(1) @@ -295,19 +325,23 @@ func renderTSRPCServiceProxies(services model.ServiceList, fullPackageName strin g.L("return") g.Ind(-1) g.L("}") + if isSessionRequest { g.Ind(-1).L("}") } + g.L("gotsrpc.Monitor(w, r, args, rets, callStats)") g.L("return") g.Ind(-1) } + g.L("default:") g.Ind(1).L("gotsrpc.ClearStats(r)") g.Ind(1).L("gotsrpc.ErrorFuncNotFound(w)") g.Ind(-2).L("}") // close switch g.Ind(-1).L("}") // close ServeHttp } + return nil } @@ -320,23 +354,32 @@ type goMethod struct { } func newMethodSignature(method *model.Method, aliases map[string]string, fullPackageName string) goMethod { - var args []string - var params []string + var ( + args []string + params []string + ) + params = append(params, "ctx go_context.Context") for _, a := range goMethodArgsWithoutHTTPContextRelatedArgs(method) { args = append(args, a.Name) params = append(params, a.Name+" "+valueGoType(a.Value, aliases, fullPackageName)) } - var rets []string - var returns []string + + var ( + rets []string + returns []string + ) + for i, r := range method.Return { name := r.Name if len(name) == 0 { name = fmt.Sprintf("ret%s_%d", method.Name, i) } + rets = append(rets, "&"+name) returns = append(returns, name+" "+valueGoType(r.Value, aliases, fullPackageName)) } + returns = append(returns, "clientErr error") return goMethod{ @@ -364,6 +407,7 @@ func renderTSRPCServiceClients(services model.ServiceList, fullPackageName strin if !config.IsTSRPC(service.Name) { continue } + for _, m := range service.Methods { extractImports(m.Args, fullPackageName, aliases) extractImports(m.Return, fullPackageName, aliases) @@ -381,6 +425,7 @@ func renderTSRPCServiceClients(services model.ServiceList, fullPackageName strin clientName := "HTTP" + interfaceName g.L(`type ` + interfaceName + ` interface { `) + for _, method := range service.Methods { ms := newMethodSignature(method, aliases, fullPackageName) g.L(ms.renderSignature()) @@ -426,6 +471,7 @@ func renderTSRPCServiceClients(services model.ServiceList, fullPackageName strin g.NL() } } + return nil } @@ -465,6 +511,7 @@ func renderGoRPCServiceProxies(services model.ServiceList, fullPackageName strin } proxyName := service.Name + "GoRPCProxy" + g.L(`type (`) g.L(` ` + proxyName + ` struct { @@ -473,30 +520,39 @@ func renderGoRPCServiceProxies(services model.ServiceList, fullPackageName strin callStatsHandler gotsrpc.GoRPCCallStatsHandlerFun } `) + for _, method := range service.Methods { g.L(ucfirst(service.Name+method.Name) + `Request struct {`) + for _, a := range goMethodArgsWithoutHTTPContextRelatedArgs(method) { g.L(ucfirst(a.Name) + ` ` + valueGoType(a.Value, aliases, fullPackageName)) } + g.L(`}`) g.L(ucfirst(service.Name+method.Name) + `Response struct {`) + for i, r := range method.Return { name := r.Name if len(name) == 0 { name = fmt.Sprintf("ret%s_%d", method.Name, i) } + g.L(ucfirst(name) + ` ` + valueGoType(r.Value, aliases, fullPackageName)) } + g.L(`}`) g.NL() } + g.L(`)`) g.NL() g.L(`func init() {`) + for _, method := range service.Methods { g.L(`gob.Register(` + ucfirst(service.Name+method.Name) + `Request{})`) g.L(`gob.Register(` + ucfirst(service.Name+method.Name) + `Response{})`) } + g.L(`}`) g.L(` func New` + proxyName + `(addr string, service ` + servicePointer + service.Name + `, tlsConfig *tls.Config) *` + proxyName + ` { @@ -538,37 +594,51 @@ func renderGoRPCServiceProxies(services model.ServiceList, fullPackageName strin g.L(`funcName := funcNameParts[len(funcNameParts)-1]`) g.NL() g.L(`switch funcName {`) + for _, method := range service.Methods { var argParams []string + nonHTTPRelatedMethodArgs := goMethodArgsWithoutHTTPContextRelatedArgs(method) + diffNONHTTPRelatedMethodArgs := len(method.Args) - len(nonHTTPRelatedMethodArgs) for i := 0; i < diffNONHTTPRelatedMethodArgs; i++ { argParams = append(argParams, "nil") } + for _, a := range nonHTTPRelatedMethodArgs { argParams = append(argParams, "req."+ucfirst(a.Name)) } - var rets []string - var retParams []string + + var ( + rets []string + retParams []string + ) + for i, r := range method.Return { name := r.Name if len(name) == 0 { name = fmt.Sprintf("ret%s_%d", method.Name, i) } + rets = append(rets, name) retParams = append(retParams, ucfirst(name)+`: `+name) } + g.L(`case "` + service.Name + method.Name + `Request":`) + if len(nonHTTPRelatedMethodArgs) > 0 { g.L(`req := request.(` + service.Name + method.Name + `Request)`) } + if len(rets) > 0 { g.L(strings.Join(rets, ", ") + ` := p.service.` + method.Name + `(` + strings.Join(argParams, ", ") + `)`) } else { g.L(`p.service.` + method.Name + `(` + strings.Join(argParams, ", ") + `)`) } + g.L(`response = ` + service.Name + method.Name + `Response{` + strings.Join(retParams, ", ") + `}`) } + g.L(`default:`) g.L(`fmt.Println("Unknown request type", reflect.TypeOf(request).String())`) g.L(`}`) @@ -585,6 +655,7 @@ func renderGoRPCServiceProxies(services model.ServiceList, fullPackageName strin g.L(`return`) g.L(`}`) } + return nil } @@ -598,6 +669,7 @@ func renderGoRPCServiceClients(services model.ServiceList, fullPackageName strin if !config.IsGoRPC(service.Name) { continue } + for _, m := range service.Methods { extractImports(m.Args, fullPackageName, aliases) extractImports(m.Return, fullPackageName, aliases) @@ -637,6 +709,7 @@ func renderGoRPCServiceClients(services model.ServiceList, fullPackageName strin } `) g.NL() + for _, method := range service.Methods { var ( args []string @@ -646,97 +719,123 @@ func renderGoRPCServiceClients(services model.ServiceList, fullPackageName strin args = append(args, ucfirst(a.Name)+`: `+a.Name) params = append(params, a.Name+" "+valueGoType(a.Value, aliases, fullPackageName)) } + var ( rets []string returns []string ) + for i, r := range method.Return { name := r.Name if len(name) == 0 { name = fmt.Sprintf("ret%s_%d", method.Name, i) } + rets = append(rets, "rpcResp."+ucfirst(name)) returns = append(returns, name+" "+valueGoType(r.Value, aliases, fullPackageName)) } + returns = append(returns, "clientErr error") g.L(`func (tsc *` + clientName + `) ` + method.Name + `(` + strings.Join(params, ", ") + `) (` + strings.Join(returns, ", ") + `) {`) g.L(`rpcReq := ` + service.Name + method.Name + `Request{` + strings.Join(args, ", ") + `}`) + if len(rets) > 0 { g.L(`rpcRes, rpcErr := tsc.Client.Call(rpcReq)`) } else { g.L(`_, rpcErr := tsc.Client.Call(rpcReq)`) } + g.L(`if rpcErr != nil {`) g.L(`clientErr = rpcErr`) g.L(`return`) g.L(`}`) + if len(rets) > 0 { g.L(`rpcResp := rpcRes.(` + service.Name + method.Name + `Response)`) g.L(`return ` + strings.Join(rets, ", ") + `, nil`) } else { g.L(`return nil`) } + g.L(`}`) g.NL() } } + return nil } func RenderGoTSRPCProxies(services model.ServiceList, longPackageName, packageName string, config *config.Target, unions map[string][]string) (gocode string, err error) { g := NewCode(" ") + err = renderTSRPCServiceProxies(services, longPackageName, packageName, config, unions, g) if err != nil { return } + gocode = g.String() + return } func RenderGoTSRPCClients(services model.ServiceList, longPackageName, packageName string, config *config.Target) (gocode string, err error) { g := NewCode(" ") + err = renderTSRPCServiceClients(services, longPackageName, packageName, config, g) if err != nil { return } + gocode = g.String() + return } func RenderGoRPCProxies(services model.ServiceList, longPackageName, packageName string, config *config.Target) (gocode string, err error) { g := NewCode(" ") + err = renderGoRPCServiceProxies(services, longPackageName, packageName, config, g) if err != nil { return } + gocode = g.String() + return } func RenderGoRPCClients(services model.ServiceList, longPackageName, packageName string, config *config.Target) (gocode string, err error) { g := NewCode(" ") + err = renderGoRPCServiceClients(services, longPackageName, packageName, config, g) if err != nil { return } + gocode = g.String() + return } func goMethodArgsWithoutHTTPContextRelatedArgs(m *model.Method) (filteredArgs []*model.Field) { filteredArgs = []*model.Field{} + for argI, arg := range m.Args { if argI == 0 && valueIsHTTPResponseWriter(arg.Value) { continue } + if argI == 1 && valueIsHTTPRequest(arg.Value) { continue } + if argI == 0 && valueIsContext(arg.Value) { continue } + filteredArgs = append(filteredArgs, arg) } + return } @@ -744,21 +843,27 @@ func renderInit(unions map[string][]string, aliases map[string]string, packageNa if len(unions) > 0 { g.L("func init() {") g.Ind(1) + var strs []string + for pkg, us := range unions { for _, name := range us { var str string if packageName != pkg && aliases[pkg] != "" { str += aliases[pkg] + "." } + str += name strs = append(strs, str) } } + sort.Strings(strs) + for _, str := range strs { g.L("gotsrpc.MustRegisterUnionExt(" + str + "{})") } + g.Ind(-1) g.L("}") } @@ -769,6 +874,7 @@ func renderImports(aliases map[string]string, packageName string) string { for importPath, alias := range aliases { imports += alias + " \"" + importPath + "\"\n" } + return ` // Code generated by gotsrpc https://github.com/foomo/gotsrpc/v2 - DO NOT EDIT. diff --git a/internal/codegen/tsclient.go b/internal/codegen/tsclient.go index be01e93..8b56c83 100644 --- a/internal/codegen/tsclient.go +++ b/internal/codegen/tsclient.go @@ -24,44 +24,61 @@ func renderTypescriptClient(service *model.Service, mappings config.TypeScriptMa for _, method := range service.Methods { ts.App("async " + lcfirst(method.Name) + "(") + var callArgs []string + argOffset := 0 + for index, arg := range method.Args { if index == 0 && valueIsHTTPResponseWriter(arg.Value) { trace("skipping first arg is a http.ResponseWriter") + argOffset = 1 + continue } else if index == 0 && valueIsContext(arg.Value) { trace("skipping first arg is a context.Context") + argOffset = 1 + continue } + if index == 1 && valueIsHTTPRequest(arg.Value) { trace("skipping second arg is a *http.Request") + argOffset = 2 + continue } } + argCount := 0 + for index, arg := range method.Args { if index < argOffset { continue } + if index > argOffset { ts.App(", ") } + ts.App(fieldTSName(arg)) ts.App(":") valueTSType(arg.Value, mappings, scalars, structs, ts, arg.JSONInfo) callArgs = append(callArgs, arg.Name) argCount++ } + ts.App("):") returnTypeTS := NewCode(" ") returnTypeTS.App("{") + innerReturnTypeTS := NewCode(" ") innerReturnTypeTS.App("{") + firstReturnType := "" countReturns := 0 countInnerReturns := 0 @@ -78,6 +95,7 @@ func renderTypescriptClient(service *model.Service, mappings config.TypeScriptMa retArgName += "_" + fmt.Sprint(index) } } + if index > 0 { returnTypeTS.App("; ") innerReturnTypeTS.App("; ") @@ -92,16 +110,24 @@ func renderTypescriptClient(service *model.Service, mappings config.TypeScriptMa valueTSType(retField.Value, mappings, scalars, structs, firstReturnTypeTS, retField.JSONInfo) firstReturnType = firstReturnTypeTS.String() } + countReturns++ + returnTypeTS.App(retArgName) returnTypeTS.App(":") + responseObject += responseObjectPrefix + retArgName + " : response[" + strconv.Itoa(index) + "]" + valueTSType(retField.Value, mappings, scalars, structs, returnTypeTS, retField.JSONInfo) + responseObjectPrefix = ", " } + responseObject += "};" + returnTypeTS.App("}") innerReturnTypeTS.App("}") + if countReturns == 0 { ts.App("Promise {") } else if countReturns == 1 { @@ -109,6 +135,7 @@ func renderTypescriptClient(service *model.Service, mappings config.TypeScriptMa } else if countReturns > 1 { ts.App("Promise<" + returnTypeTS.String() + "> {") } + ts.NL() ts.Ind(1) @@ -119,6 +146,7 @@ func renderTypescriptClient(service *model.Service, mappings config.TypeScriptMa } call := "this.transport<" + innerCallTypeString + ">(\"" + method.Name + "\", [" + strings.Join(callArgs, ", ") + "])" + switch countReturns { case 0: ts.L("await " + call) @@ -133,7 +161,9 @@ func renderTypescriptClient(service *model.Service, mappings config.TypeScriptMa ts.App("}") ts.NL() } + ts.Ind(-1) ts.L("}") + return nil } diff --git a/internal/codegen/typescript.go b/internal/codegen/typescript.go index 78e582e..aa24adb 100644 --- a/internal/codegen/typescript.go +++ b/internal/codegen/typescript.go @@ -17,6 +17,7 @@ func fieldTSName(f *model.Field) string { if f.JSONInfo != nil && len(f.JSONInfo.Name) > 0 { n = f.JSONInfo.Name } + return n } @@ -33,18 +34,23 @@ func valueTSType(v *model.Value, mappings config.TypeScriptMappings, scalars map if enumKey { ts.App("Partial<") } + ts.App("Record<") + if v.Map.Key != nil { valueTSType(v.Map.Key, mappings, scalars, structs, ts, nil) } else { ts.App(v.Map.KeyType) } + ts.App(",") valueTSType(v.Map.Value, mappings, scalars, structs, ts, nil) ts.App(">") + if enumKey { ts.App(">") } + if jsonInfo == nil || !jsonInfo.OmitEmpty { ts.App("|null") } @@ -58,53 +64,67 @@ func valueTSType(v *model.Value, mappings config.TypeScriptMappings, scalars map } else { valueTSType(v.Array.Value, mappings, scalars, structs, ts, nil) } + if jsonInfo == nil || !jsonInfo.OmitEmpty { ts.App("|null") } case v.Scalar != nil: if v.Scalar.Package != "" { mapping, ok := mappings[v.Scalar.Package] + var tsModule string if ok { tsModule = mapping.TypeScriptModule } + tsType := tsModule + "." + tsTypeFromScalarType(v.ScalarType) if value, ok := tsTypeAliases[tsType]; ok { tsType = value } + ts.App(tsType) + if v.IsPtr && (jsonInfo == nil || !jsonInfo.OmitEmpty) { ts.App("|null") } + return } + ts.App(tsTypeFromScalarType(v.Scalar.Type)) case v.StructType != nil: if v.StructType.Package != "" { mapping, ok := mappings[v.StructType.Package] + var tsModule string if ok { tsModule = mapping.TypeScriptModule } + ts.App(tsModule + "." + v.StructType.Name) + hiddenStruct, isHiddenStruct := structs[v.StructType.FullName()] if isHiddenStruct && (hiddenStruct.Array != nil || hiddenStruct.Map != nil) && (jsonInfo == nil || !jsonInfo.OmitEmpty) { ts.App("|null") } else if v.IsPtr && (jsonInfo == nil || !jsonInfo.OmitEmpty) { ts.App("|null") } + return } + ts.App(v.StructType.Name) case v.Struct != nil: ts.L("{").Ind(1) renderStructFields(v.Struct.Fields, mappings, scalars, structs, ts) ts.Ind(-1).App("}") + if v.IsPtr && (jsonInfo == nil || !jsonInfo.OmitEmpty) { ts.App("|null") } case len(v.ScalarType) > 0: ts.App(tsTypeFromScalarType(v.ScalarType)) + if v.IsPtr && (jsonInfo == nil || !jsonInfo.OmitEmpty) { ts.App("|null") } @@ -122,6 +142,7 @@ func tsTypeFromScalarType(scalarType model.ScalarType) string { case model.ScalarTypeBool: return "boolean" } + return string(scalarType) } @@ -132,10 +153,13 @@ func renderStructFields(fields []*model.Field, mappings config.TypeScriptMapping } else if f.JSONInfo != nil && f.JSONInfo.Ignore { continue } + ts.App(fieldTSName(f)) + if f.JSONInfo != nil && f.JSONInfo.OmitEmpty { ts.App("?") } + ts.App(":") valueTSType(f.Value, mappings, scalars, structs, ts, f.JSONInfo) ts.App(";") @@ -145,6 +169,7 @@ func renderStructFields(fields []*model.Field, mappings config.TypeScriptMapping func renderTypescriptStruct(str *model.Struct, mappings config.TypeScriptMappings, scalars map[string]*model.Scalar, structs map[string]*model.Struct, ts *Code) error { ts.L("// " + str.FullName()) + switch { case str.Array != nil: if str.Array.Len > 0 && str.Array.Value.ScalarType == model.ScalarTypeByte { @@ -154,6 +179,7 @@ func renderTypescriptStruct(str *model.Struct, mappings config.TypeScriptMapping valueTSType(str.Array.Value, mappings, scalars, structs, ts, nil) ts.App(">") } + ts.NL() case str.Map != nil: enumKey := str.Map.Key != nil && str.Map.Key.Scalar != nil @@ -162,113 +188,145 @@ func renderTypescriptStruct(str *model.Struct, mappings config.TypeScriptMapping } else { ts.App("export type " + str.Name + " = Record<") } + if str.Map.Key != nil { valueTSType(str.Map.Key, mappings, scalars, structs, ts, nil) } else { ts.App(str.Map.KeyType) } + ts.App(",") valueTSType(str.Map.Value, mappings, scalars, structs, ts, nil) + if enumKey { ts.App(">>") } else { ts.App(">") } + ts.NL() // special handling of inline only structs case len(str.UnionFields) > 0: if len(str.Fields) > 0 || len(str.InlineFields) > 0 { return errors.New("no fields or inline fields are supported when using union") } + switch { case str.UnionFields[0].Value.StructType != nil: ts.App("export type " + str.Name + " = ") + var isUndefined bool + for i, unionField := range str.UnionFields { if i > 0 { ts.App(" | ") } + valueTSType(unionField.Value, mappings, scalars, structs, ts, &model.JSONInfo{OmitEmpty: true}) + if unionField.Value.IsPtr { isUndefined = true } } + if isUndefined { ts.App(" | undefined") } + ts.NL() case str.UnionFields[0].Value.Scalar != nil: ts.App("export const " + str.Name + " = ") ts.App("{ ") + for i, field := range str.UnionFields { if i > 0 { ts.App(", ") } + ts.App("...") valueTSType(field.Value, mappings, scalars, structs, ts, &model.JSONInfo{OmitEmpty: true}) } + ts.App(" }") ts.NL() ts.App("export type " + str.Name + " = ") + for i, field := range str.UnionFields { if i > 0 { ts.App(" | ") } + valueTSType(field.Value, mappings, scalars, structs, ts, &model.JSONInfo{OmitEmpty: true}) } + ts.NL() default: return errors.New("could not resolve this union type") } case len(str.InlineFields) > 0: var extends bool + ts.App("export interface " + str.Name) + for i, inlineField := range str.InlineFields { if inlineField.Value.Scalar != nil { if _, isStruct := structs[inlineField.Value.Scalar.Package+"."+inlineField.Value.Scalar.Name]; !isStruct { continue } } + if !extends { ts.App(" extends ") + extends = true } + if i > 0 { ts.App(", ") } + if inlineField.Value.IsPtr { ts.App("Partial<") } + valueTSType(inlineField.Value, mappings, scalars, structs, ts, &model.JSONInfo{OmitEmpty: true}) + if inlineField.Value.IsPtr { ts.App(">") } } + ts.App(" {") ts.NL() ts.Ind(1) + for _, inlineField := range str.InlineFields { if inlineField.Value.Scalar != nil { if _, isStruct := structs[inlineField.Value.Scalar.Package+"."+inlineField.Value.Scalar.Name]; isStruct { continue // already handled as extends } + if inlineField.JSONInfo != nil && inlineField.JSONInfo.Ignore { continue } + if n := fieldTSName(inlineField); n != "" { ts.App(n) } else { ts.App(inlineField.Value.Scalar.Name) } + if inlineField.JSONInfo != nil && inlineField.JSONInfo.OmitEmpty { ts.App("?") } + ts.App(":") valueTSType(inlineField.Value, mappings, scalars, structs, ts, &model.JSONInfo{OmitEmpty: true}) ts.App(";") ts.NL() } } + renderStructFields(str.Fields, mappings, scalars, structs, ts) ts.Ind(-1).L("}") default: @@ -276,6 +334,7 @@ func renderTypescriptStruct(str *model.Struct, mappings config.TypeScriptMapping renderStructFields(str.Fields, mappings, scalars, structs, ts) ts.Ind(-1).L("}") } + return nil } @@ -290,17 +349,21 @@ func RenderTypescriptStructsToPackages( for _, mapping := range mappings { codeMap[mapping.GoPackage] = map[string]*Code{} } + for name, str := range structs { if str == nil { err = errors.New("could not resolve: " + name) return } + packageCodeMap, ok := codeMap[str.Package] if !ok { err = errors.New("missing code mapping for go package : " + str.Package + " => you have to add a mapping from this go package to a TypeScript module in your build-config.yml in the mappings section") return } + packageCodeMap[str.Name] = NewCode(" ") + err = renderTypescriptStruct(str, mappings, scalarTypes, structs, packageCodeMap[str.Name]) if err != nil { return @@ -314,6 +377,7 @@ func RenderTypescriptStructsToPackages( err = errors.New("missing code mapping for go package : " + packageName + " => you have to add a mapping from this go package to a TypeScript module in your build-config.yml in the mappings section") return } + for packageConstantTypeName, packageConstantTypeValues := range packageConstantTypes { packageCodeMap[packageConstantTypeName] = NewCode(" ") packageCodeMap[packageConstantTypeName].L("// " + packageName + "." + packageConstantTypeName) @@ -323,12 +387,15 @@ func RenderTypescriptStructsToPackages( for k := range packageConstantTypeValuesList { keys = append(keys, k) } + sort.Strings(keys) packageCodeMap[packageConstantTypeName].L("export enum " + packageConstantTypeName + " {").Ind(1) + for _, k := range keys { enum := strings.TrimPrefix(stringsx.ToCamel(k), packageConstantTypeName) packageCodeMap[packageConstantTypeName].L(enum + " = " + packageConstantTypeValuesList[k].Value + ",") } + packageCodeMap[packageConstantTypeName].Ind(-1).L("}") } else if packageConstantTypeValuesString, ok := packageConstantTypeValues.(string); ok { packageCodeMap[packageConstantTypeName].L("export type " + packageConstantTypeName + " = " + packageConstantTypeValuesString) @@ -343,41 +410,51 @@ func RenderTypescriptStructsToPackages( mappedTypeScript[goPackage] = map[string]*Code{} } } + for _, mapping := range mappings { for structName, structCode := range codeMap[mapping.GoPackage] { ensureCodeInPackage(mapping.GoPackage) mappedTypeScript[mapping.GoPackage][structName] = structCode } } + return nil } func Split(str string, seps []string) []string { var res []string + strs := []string{str} + for _, sep := range seps { var nextStrs []string for _, str := range strs { nextStrs = append(nextStrs, strings.Split(str, sep)...) } + strs = nextStrs res = nextStrs } + return res } func RenderTypeScriptServices(services model.ServiceList, mappings config.TypeScriptMappings, scalars map[string]*model.Scalar, structs map[string]*model.Struct, target *config.Target) (typeScript string, err error) { ts := NewCode(" ") + for _, service := range services { if !target.IsTSRPC(service.Name) { continue } + err = renderTypescriptClient(service, mappings, scalars, structs, ts) if err != nil { return } } + typeScript = ts.String() + return } diff --git a/internal/model/scalar.go b/internal/model/scalar.go index 56e9c7b..f1bd32a 100644 --- a/internal/model/scalar.go +++ b/internal/model/scalar.go @@ -11,5 +11,6 @@ func (st *Scalar) FullName() string { if len(fullName) == 0 { fullName = st.Name } + return fullName } diff --git a/internal/model/struct.go b/internal/model/struct.go index 63f96eb..13b1a2f 100644 --- a/internal/model/struct.go +++ b/internal/model/struct.go @@ -16,5 +16,6 @@ func (s *Struct) FullName() string { if len(fullName) == 0 { fullName = s.Name } + return fullName } diff --git a/internal/model/structtype.go b/internal/model/structtype.go index efdff89..a4af6be 100644 --- a/internal/model/structtype.go +++ b/internal/model/structtype.go @@ -10,5 +10,6 @@ func (st *StructType) FullName() string { if len(fullName) == 0 { fullName = st.Name } + return fullName } diff --git a/internal/parser/fileimportspecmap.go b/internal/parser/fileimportspecmap.go index 18c7c28..fd34b1a 100644 --- a/internal/parser/fileimportspecmap.go +++ b/internal/parser/fileimportspecmap.go @@ -7,5 +7,6 @@ func (m fileImportSpecMap) getPackagePath(packageName string) string { if ok { packageName = is.path } + return packageName } diff --git a/internal/parser/parser.go b/internal/parser/parser.go index 6925de4..bb50e34 100644 --- a/internal/parser/parser.go +++ b/internal/parser/parser.go @@ -25,13 +25,17 @@ func parseDir(goPaths []string, gomod config.Namespace, packageName string) (map fset := token.NewFileSet() dir := strings.Replace(packageName, gomod.Name, gomod.Path, 1) pkgs, err := parser.ParseDir(fset, dir, parserExcludeFiles, parser.DeclarationErrors|parser.AllErrors) + return pkgs, fset, err } errorStrings := map[string]string{} + for _, goPath := range goPaths { var dir string + fset := token.NewFileSet() + if gomod.ModFile != nil { for _, rep := range gomod.ModFile.Replace { if packageName == rep.Old.Path || strings.HasPrefix(packageName, rep.Old.Path+"/") { @@ -42,19 +46,23 @@ func parseDir(goPaths []string, gomod config.Namespace, packageName string) (map trace("replacing package", packageName, rep.Old.String(), rep.New.String()) dir = strings.TrimSuffix(path.Join(goPath, "pkg", "mod", rep.New.String(), strings.TrimPrefix(packageName, rep.Old.Path)), "/") } + break } } + if dir == "" { for _, req := range gomod.ModFile.Require { if packageName == req.Mod.Path || strings.HasPrefix(packageName, req.Mod.Path+"/") { trace("resolving mod package", packageName, req.Mod.String()) dir = strings.TrimSuffix(path.Join(goPath, "pkg", "mod", req.Mod.String(), strings.TrimPrefix(packageName, req.Mod.Path)), "/") + break } } } } + if dir == "" { if strings.HasSuffix(goPath, "vendor") { dir = path.Join(goPath, packageName) @@ -62,12 +70,15 @@ func parseDir(goPaths []string, gomod config.Namespace, packageName string) (map dir = path.Join(goPath, "src", packageName) } } + pkgs, err := parser.ParseDir(fset, dir, parserExcludeFiles, parser.AllErrors) if err == nil { return pkgs, fset, nil } + errorStrings[dir] = err.Error() } + return nil, nil, errors.New("could not parse dir for package name: " + packageName + " in goPaths " + strings.Join(goPaths, ", ") + " : " + fmt.Sprint(errorStrings)) } @@ -76,10 +87,12 @@ func parsePackage(goPaths []string, gomod config.Namespace, packageName string) if err != nil { return nil, errors.New("could not parse package " + packageName + ": " + err.Error()) } + packageNameParts := strings.Split(packageName, "/") if len(packageNameParts) == 0 { return nil, errors.New("invalid package name given") } + strippedPackageName := packageNameParts[len(packageNameParts)-1] if len(pkgs) == 1 { for _, v := range pkgs { @@ -87,7 +100,9 @@ func parsePackage(goPaths []string, gomod config.Namespace, packageName string) break } } + var foundPackages []string + sortedGoPaths := make([]string, len(goPaths)) copy(sortedGoPaths, goPaths) sort.Sort(byLen(sortedGoPaths)) @@ -106,6 +121,7 @@ Loop: prefix := goPath + "/" if strings.HasPrefix(pkgFile, prefix) && !strings.HasSuffix(pkgFile, "_test.go") && !strings.HasSuffix(pkgFile, "_generator.go") { trimmedFilename := strings.TrimPrefix(pkgFile, prefix) + parts := strings.Split(trimmedFilename, "/") if len(parts) > 1 { parts = parts[0 : len(parts)-1] @@ -127,5 +143,6 @@ Loop: // create new package with resolved objects resolvedPkg, _ := ast.NewPackage(fset, parsedPkg.Files, nil, nil) // ignore error + return resolvedPkg, nil } diff --git a/internal/parser/servicereader.go b/internal/parser/servicereader.go index c560526..5046681 100644 --- a/internal/parser/servicereader.go +++ b/internal/parser/servicereader.go @@ -32,12 +32,15 @@ func Read( err = errors.New("nothing to do service names are empty") return } + pkg, parseErr := parsePackage(goPaths, gomod, packageName) if parseErr != nil { err = parseErr return } + pkgName = pkg.Name + services, err = readServicesInPackage(pkg, packageName, serviceMap) if err != nil { return @@ -51,6 +54,7 @@ func Read( collectScalarTypes(m.Args, missingTypes) } } + trace("missing") traceData(missingTypes) @@ -61,13 +65,16 @@ func Read( if collectErr != nil { err = errors.New("error while collecting structs: " + collectErr.Error()) } + trace("---------------- found structs -------------------") traceData(structs) trace("---------------- /found structs -------------------") trace("---------------- found scalars -------------------") traceData(scalars) trace("---------------- /found scalars -------------------") + allConstantTypes := map[string]map[string]interface{}{} + for _, structDef := range structs { if structDef != nil { structPackage := structDef.Package @@ -81,6 +88,7 @@ func Read( } } } + for _, scalarDef := range scalars { if scalarDef != nil { scalarPackage := scalarDef.Package @@ -101,6 +109,7 @@ func Read( } constantTypes = map[string]map[string]interface{}{} + for constantTypePackage, constantType := range allConstantTypes { for constantTypeName, constantTypeVales := range constantType { fullName := constantTypePackage + "." + constantTypeName @@ -110,9 +119,11 @@ func Read( if scalarOK || structOK || constantsOK { missingConstants[fullName] = false + if _, ok := constantTypes[constantTypePackage]; !ok { constantTypes[constantTypePackage] = map[string]interface{}{} } + constantTypes[constantTypePackage][constantTypeName] = constantTypeVales } } @@ -132,7 +143,9 @@ func Read( fixFieldStructs(method.Return, structs, scalars) } } + traceData("---------------------------", services) + return } @@ -143,6 +156,7 @@ func readServiceFile(file *ast.File, packageName string, services model.ServiceL return service, true } } + return nil, false } @@ -152,11 +166,13 @@ func readServiceFile(file *ast.File, packageName string, services model.ServiceL if funcDecl, ok := decl.(*ast.FuncDecl); ok { if funcDecl.Recv != nil { trace("that is a method named", funcDecl.Name) + if len(funcDecl.Recv.List) == 1 { firstReceiverField := funcDecl.Recv.List[0] if starExpr, ok := firstReceiverField.Type.(*ast.StarExpr); ok { if ident, ok := starExpr.X.(*ast.Ident); ok { service, ok := findService(ident.Name) + firstCharOfMethodName := funcDecl.Name.Name[0:1] if !ok || strings.ToLower(firstCharOfMethodName) == firstCharOfMethodName { continue @@ -179,18 +195,22 @@ func readServiceFile(file *ast.File, packageName string, services model.ServiceL if genDecl.Tok != token.TYPE { continue } + for _, spec := range genDecl.Specs { if typeSpec, ok := spec.(*ast.TypeSpec); ok { ident := typeSpec.Name trace("that is an interface named", ident.Name) + if service, ok := findService(ident.Name); ok { if iSpec, ok := typeSpec.Type.(*ast.InterfaceType); ok { service.IsInterface = true + for _, fieldDecl := range iSpec.Methods.List { if funcDecl, ok := fieldDecl.Type.(*ast.FuncType); ok { if len(fieldDecl.Names) == 0 { continue } + mname := fieldDecl.Names[0] trace(" on sth:", mname.Name) service.Methods = append(service.Methods, &model.Method{ @@ -206,15 +226,19 @@ func readServiceFile(file *ast.File, packageName string, services model.ServiceL } } } + for _, s := range services { sort.Sort(s.Methods) } + return nil } func readFields(fieldList *ast.FieldList, fileImports fileImportSpecMap) (fields []*model.Field) { trace("reading fields") + fields = []*model.Field{} + if fieldList == nil { return } @@ -228,7 +252,9 @@ func readFields(fieldList *ast.FieldList, fileImports fileImportSpecMap) (fields }) } } + trace("done reading fields") + return } @@ -236,6 +262,7 @@ func readServicesInPackage(pkg *ast.Package, packageName string, serviceMap map[ if pkg == nil { return nil, errors.New("package cannot be nil") } + services = model.ServiceList{} for endpoint, serviceName := range serviceMap { services = append(services, &model.Service{ @@ -244,36 +271,44 @@ func readServicesInPackage(pkg *ast.Package, packageName string, serviceMap map[ Endpoint: endpoint, }) } + pkgFiles := make([]string, 0, len(pkg.Files)) for k := range pkg.Files { pkgFiles = append(pkgFiles, k) } + sort.Strings(pkgFiles) for _, k := range pkgFiles { file := pkg.Files[k] + err = readServiceFile(file, packageName, services) if err != nil { return } } + sort.Sort(services) + return } func loadConstantTypes(pkg *ast.Package) map[string]interface{} { constantTypes := map[string]interface{}{} + for _, file := range pkg.Files { for _, decl := range file.Decls { if genDecl, ok := decl.(*ast.GenDecl); ok { switch genDecl.Tok { //nolint:exhaustive case token.TYPE: trace("got a type", genDecl.Specs) + for _, spec := range genDecl.Specs { if spec, ok := spec.(*ast.TypeSpec); ok { if _, ok := constantTypes[spec.Name.Name]; ok { continue } + switch specType := spec.Type.(type) { case *ast.InterfaceType: constantTypes[spec.Name.Name] = "any" @@ -299,6 +334,7 @@ func loadConstantTypes(pkg *ast.Package) map[string]interface{} { } case token.CONST: trace("got a const", genDecl.Specs) + for _, spec := range genDecl.Specs { if spec, ok := spec.(*ast.ValueSpec); ok { if specType, ok := spec.Type.(*ast.Ident); ok { @@ -309,6 +345,7 @@ func loadConstantTypes(pkg *ast.Package) map[string]interface{} { } else if _, ok := constantTypes[specType.Name].(map[string]*ast.BasicLit); !ok { constantTypes[specType.Name] = map[string]*ast.BasicLit{} } + constantTypes[specType.Name].(map[string]*ast.BasicLit)[spec.Names[0].Name] = valType //nolint:forcetypeassert } } @@ -319,6 +356,7 @@ func loadConstantTypes(pkg *ast.Package) map[string]interface{} { } } } + return constantTypes } @@ -327,15 +365,18 @@ func loadFlatStructs(s *model.Struct, flatStructs map[string]bool) { if s.Map.Key != nil { loadFlatStructsValue(s.Map.Key, flatStructs) } + if s.Map.Value != nil && s.Map.Value.Scalar != nil { loadFlatStructsValue(s.Map.Value, flatStructs) } } + if s.Fields != nil { for _, field := range s.Fields { loadFlatStructsValue(field.Value, flatStructs) } } + flatStructs[s.FullName()] = true } @@ -344,13 +385,16 @@ func loadFlatStructsValue(s *model.Value, flatStructs map[string]bool) { if s.Map.Key != nil { loadFlatStructsValue(s.Map.Key, flatStructs) } + if s.Map.Value != nil && s.Map.Value.Scalar != nil { loadFlatStructsValue(s.Map.Value, flatStructs) } } + if s.Struct != nil { loadFlatStructs(s.Struct, flatStructs) } + if s.Scalar != nil { flatStructs[s.Scalar.FullName()] = true } @@ -360,11 +404,13 @@ func fixFieldStructs(fields []*model.Field, structs map[string]*model.Struct, sc for _, f := range fields { if f.Value.StructType != nil { name := f.Value.StructType.FullName() + s, strctExists := structs[name] if strctExists { f.Value.IsError = s.IsError continue } + scalar, scalarExists := scalars[name] if scalarExists { f.Value.StructType = nil @@ -379,21 +425,25 @@ func collectTypes(goPaths []string, gomod config.Namespace, missingTypes map[str scannedPackageScalars := map[string]map[string]*model.Scalar{} missingTypeNames := func() []string { var missing []string + for name, isMissing := range missingTypes { if isMissing { missing = append(missing, name) } } + return missing } lastNumMissing := len(missingTypeNames()) for typesPending(structs, scalars, missingTypes) { trace("pending", missingTypeNames()) + for fullName, typeIsMissing := range missingTypes { if !typeIsMissing { continue } + fullNameParts := strings.Split(fullName, ".") fullNameParts = fullNameParts[:len(fullNameParts)-1] @@ -402,6 +452,7 @@ func collectTypes(goPaths []string, gomod config.Namespace, missingTypes map[str trace(fullName, "==========================>", fullNameParts, "=============>", packageName) packageStructs, structOK := scannedPackageStructs[packageName] + packageScalars, scalarOK := scannedPackageScalars[packageName] if !structOK || !scalarOK { parsedPackageStructs, parsedPackageScalars, err := getTypesInPackage(goPaths, gomod, packageName) @@ -410,36 +461,47 @@ func collectTypes(goPaths []string, gomod config.Namespace, missingTypes map[str } trace("found structs in", goPaths, packageName) + for structName, strct := range packageStructs { trace(" struct", structName, strct) + if strct == nil { panic("how could that be") } } + trace("found scalars in", goPaths, packageName) + for scalarName, scalar := range packageScalars { trace(" scalar", scalarName, scalar) } + traceData(parsedPackageScalars) + packageStructs = parsedPackageStructs packageScalars = parsedPackageScalars scannedPackageStructs[packageName] = packageStructs scannedPackageScalars[packageName] = packageScalars } + traceData("packageStructs", packageName, packageStructs) + for packageStructName, packageStruct := range packageStructs { missing, needed := missingTypes[packageStructName] if needed && missing { trace("picked up package struct", packageStructName, packageStruct) missingTypes[packageStructName] = false + if packageStruct == nil { panic("waaaaaaaaa") } + structs[packageStructName] = packageStruct } } traceData("packageScalars", packageScalars) + for packageScalarName, packageScalar := range packageScalars { missing, needed := missingTypes[packageScalarName] if needed && missing { @@ -449,24 +511,31 @@ func collectTypes(goPaths []string, gomod config.Namespace, missingTypes map[str } } } + newNumMissingTypes := len(missingTypeNames()) if newNumMissingTypes > 0 && newNumMissingTypes == lastNumMissing { for scalarName, scalars := range scannedPackageScalars { fmt.Println("scanned scalars ", scalarName) + for _, scalar := range scalars { fmt.Println(" ", scalar.Name) } } + for structName, strcts := range scannedPackageStructs { fmt.Println("scanned struct ", structName) + for _, strct := range strcts { fmt.Println(" ", strct.Name) } } + return errors.New(fmt.Sprintln("could not resolve at least one of the following types", missingTypeNames())) } + lastNumMissing = newNumMissingTypes } + return nil } @@ -476,11 +545,13 @@ func typesPending(structs map[string]*model.Struct, scalars map[string]*model.Sc return true } } + for _, structType := range structs { if !depsSatisfied(structType, missingTypes, structs, scalars) { return true } } + return false } @@ -503,22 +574,27 @@ func needsWorkValue(value *model.Value, needsWork func(fullName string) bool) bo return true } } + return false } func depsSatisfied(s *model.Struct, missingTypes map[string]bool, structs map[string]*model.Struct, scalars map[string]*model.Scalar) bool { needsWork := func(fullName string) bool { strct, strctOK := structs[fullName] + scalar, scalarOK := scalars[fullName] if !strctOK && !scalarOK { missingTypes[fullName] = true trace("need work ----------------------" + fullName) + return true } + if strct == nil && scalar == nil { trace("need work ----------------------" + fullName) return true } + return false } @@ -528,6 +604,7 @@ func depsSatisfied(s *model.Struct, missingTypes map[string]bool, structs map[st return false } } + return true } if ok := needWorksFields(s.Fields); !ok { @@ -537,19 +614,23 @@ func depsSatisfied(s *model.Struct, missingTypes map[string]bool, structs map[st } else if ok := needWorksFields(s.UnionFields); !ok { return false } + if s.Array != nil { if s.Array.Value != nil && needsWorkValue(s.Array.Value, needsWork) { return false } } + if s.Map != nil { if s.Map.Key != nil && needsWorkValue(s.Map.Key, needsWork) { return false } + if s.Map.Value != nil && needsWorkValue(s.Map.Value, needsWork) { return false } } + return !needsWork(s.FullName()) } @@ -562,15 +643,18 @@ func getTypesInPackage(goPaths []string, gomod config.Namespace, packageName str if err != nil { return nil, nil, err } + structs, scalars, err = readStructs(pkg, packageName) if err != nil { return nil, nil, err } + return structs, scalars, nil } func getStructTypeForField(value *model.Value) *model.StructType { var strType *model.StructType + switch { case value.StructType != nil: strType = value.StructType @@ -579,11 +663,13 @@ func getStructTypeForField(value *model.Value) *model.StructType { case value.Array != nil: strType = getStructTypeForField(value.Array.Value) } + return strType } func getScalarForField(value *model.Value) []*model.Scalar { var scalarTypes []*model.Scalar + switch { case value.Scalar != nil: scalarTypes = append(scalarTypes, value.Scalar) @@ -593,10 +679,12 @@ func getScalarForField(value *model.Value) []*model.Scalar { scalarTypes = append(scalarTypes, v...) } } + scalarTypes = append(scalarTypes, getScalarForField(value.Map.Value)...) case value.Array != nil: scalarTypes = append(scalarTypes, getScalarForField(value.Array.Value)...) } + return scalarTypes } @@ -608,6 +696,7 @@ func collectScalarTypes(fields []*model.Field, scalarTypes map[string]bool) { if len(scalarType.Package) == 0 { fullName = scalarType.Name } + switch fullName { case "error", "net/http.Request", "net/http.ResponseWriter", "context.Context": continue @@ -627,6 +716,7 @@ func collectStructTypes(fields []*model.Field, structTypes map[string]bool) { if len(strType.Package) == 0 { fullName = strType.Name } + switch fullName { case "error", "net/http.Request", "net/http.ResponseWriter", "context.Context": continue diff --git a/internal/parser/trace.go b/internal/parser/trace.go index ad03a83..d2f2840 100644 --- a/internal/parser/trace.go +++ b/internal/parser/trace.go @@ -23,6 +23,7 @@ func traceData(args ...interface{}) { trace(arg) continue } + trace(string(yamlBytes)) } } diff --git a/internal/parser/typereader.go b/internal/parser/typereader.go index db689b6..f100a9f 100644 --- a/internal/parser/typereader.go +++ b/internal/parser/typereader.go @@ -18,17 +18,21 @@ func standardImportName(importPath string) string { func getFileImports(file *ast.File, packageName string) (imports fileImportSpecMap) { imports = fileImportSpecMap{"": importSpec{alias: "", name: "", path: packageName}} + for _, decl := range file.Decls { if genDecl, ok := decl.(*ast.GenDecl); ok { if genDecl.Tok == token.IMPORT { trace("got an import", genDecl.Specs) + for _, spec := range genDecl.Specs { if spec, ok := spec.(*ast.ImportSpec); ok { importPath := spec.Path.Value[1 : len(spec.Path.Value)-1] + importName := spec.Name.String() if importName == "" || importName == "" { importName = standardImportName(importPath) } + imports[importName] = importSpec{ alias: importName, name: standardImportName(importPath), @@ -39,6 +43,7 @@ func getFileImports(file *ast.File, packageName string) (imports fileImportSpecM } } } + return imports } @@ -46,6 +51,7 @@ func extractJSONInfo(tag string) *model.JSONInfo { structTag := reflect.StructTag(tag) jsonTags := strings.Split(structTag.Get("json"), ",") + gotsrpcTags := strings.Split(structTag.Get("gotsrpc"), ",") if len(jsonTags) == 0 && len(gotsrpcTags) == 0 { return nil @@ -54,6 +60,7 @@ func extractJSONInfo(tag string) *model.JSONInfo { for k, value := range jsonTags { jsonTags[k] = strings.TrimSpace(value) } + for k, value := range gotsrpcTags { gotsrpcTags[k] = strings.TrimSpace(value) } @@ -75,6 +82,7 @@ func extractJSONInfo(tag string) *model.JSONInfo { name = jsonTags[0] } } + if len(jsonTags) > 1 { for _, value := range jsonTags[1:] { switch value { @@ -129,6 +137,7 @@ func getScalarFromAstIdent(ident *ast.Ident) model.ScalarType { } else if ident.Obj == nil { return model.ScalarType(ident.Name) } + return model.ScalarTypeNone } } @@ -139,11 +148,13 @@ func getTypesFromAstType(ident *ast.Ident) (structType string, scalarType model. case model.ScalarTypeNone: structType = ident.Name } + return } func readAstType(v *model.Value, fieldIdent *ast.Ident, fileImports fileImportSpecMap, packageName string) { structType, scalarType := getTypesFromAstType(fieldIdent) + v.ScalarType = scalarType if len(structType) > 0 { v.StructType = &model.StructType{ @@ -166,6 +177,7 @@ func readAstType(v *model.Value, fieldIdent *ast.Ident, fileImports fileImportSp func readAstStarExpr(v *model.Value, starExpr *ast.StarExpr, fileImports fileImportSpecMap) { v.IsPtr = true + switch starExprType := starExpr.X.(type) { case *ast.Ident: readAstType(v, starExprType, fileImports, "") @@ -181,6 +193,7 @@ func readAstStarExpr(v *model.Value, starExpr *ast.StarExpr, fileImports fileImp func readAstMapType(m *model.Map, mapType *ast.MapType, fileImports fileImportSpecMap) { trace(" map key", mapType.Key, reflect.ValueOf(mapType.Key).Type().String()) trace(" map value", mapType.Value, reflect.ValueOf(mapType.Value).Type().String()) + switch keyType := mapType.Key.(type) { case *ast.Ident: _, scalarType := getTypesFromAstType(keyType) @@ -193,6 +206,7 @@ func readAstMapType(m *model.Map, mapType *ast.MapType, fileImports fileImportSp readAstSelectorExpr(m.Key, keyType, fileImports) default: } + loadValueExpr(m.Value, mapType.Value, fileImports) } @@ -200,6 +214,7 @@ func readAstSelectorExpr(v *model.Value, selectorExpr *ast.SelectorExpr, fileImp switch selExpType := selectorExpr.X.(type) { case *ast.Ident: readAstType(v, selectorExpr.Sel, fileImports, selExpType.Name) + if v.StructType != nil { v.StructType.Package = fileImports.getPackagePath(v.StructType.Name) v.StructType.Name = selectorExpr.Sel.Name @@ -222,6 +237,7 @@ func loadValueExpr(v *model.Value, expr ast.Expr, fileImports fileImportSpecMap) switch exprType := expr.(type) { case *ast.ArrayType: v.Array = &model.Array{Value: &model.Value{}} + if exprType.Len != nil { if lit, ok := exprType.Len.(*ast.BasicLit); ok { if n, err := strconv.Atoi(lit.Value); err == nil { @@ -279,16 +295,20 @@ func readField(astField *ast.Field, fileImports fileImportSpecMap) (names []stri names = append(names, name.Name) } } + v = &model.Value{} loadValueExpr(v, astField.Type, fileImports) + if astField.Tag != nil { jsonInfo = extractJSONInfo(astField.Tag.Value[1 : len(astField.Tag.Value)-1]) } + return } func readFieldList(fieldList []*ast.Field, fileImports fileImportSpecMap) (fields []*model.Field, inlineFields []*model.Field, unionFields []*model.Field) { fields = []*model.Field{} + for _, field := range fieldList { if names, value, jsonInfo := readField(field, fileImports); value != nil { for _, name := range names { @@ -305,6 +325,7 @@ func readFieldList(fieldList []*ast.Field, fileImports fileImportSpecMap) (field Value: value, JSONInfo: jsonInfo, }) + continue } } else if strings.Compare(strings.ToLower(name[:1]), name[:1]) == 0 { @@ -315,8 +336,10 @@ func readFieldList(fieldList []*ast.Field, fileImports fileImportSpecMap) (field Value: value, JSONInfo: jsonInfo, }) + continue } + fields = append(fields, &model.Field{ Name: name, Value: value, @@ -325,6 +348,7 @@ func readFieldList(fieldList []*ast.Field, fileImports fileImportSpecMap) (field } } } + return } @@ -348,6 +372,7 @@ func extractErrorTypes(file *ast.File, packageName string, errorTypes map[string } } } + return } @@ -366,6 +391,7 @@ func extractTypes(file *ast.File, packageName string, structs map[string]*model. Package: packageName, } trace("StructType", obj.Name) + fields, inlineFields, unionFields := readFieldList(typeSpecType.Fields.List, fileImports) structs[structName].Fields = fields structs[structName].InlineFields = inlineFields @@ -412,14 +438,18 @@ func extractTypes(file *ast.File, packageName string, structs map[string]*model. } } } + return nil } func readStructs(pkg *ast.Package, packageName string) (structs map[string]*model.Struct, scalars map[string]*model.Scalar, err error) { structs = map[string]*model.Struct{} + trace("reading files in package", packageName) + scalars = map[string]*model.Scalar{} errorTypes := map[string]bool{} + for _, file := range pkg.Files { err = extractTypes(file, packageName, structs, scalars) if err != nil { @@ -431,11 +461,13 @@ func readStructs(pkg *ast.Package, packageName string) (structs map[string]*mode return } } + for name, structType := range structs { _, isErrorType := errorTypes[name] if isErrorType { structType.IsError = true } } + return } diff --git a/responsewriter.go b/responsewriter.go index 2fca722..41000c4 100644 --- a/responsewriter.go +++ b/responsewriter.go @@ -20,5 +20,6 @@ func (r *ResponseWriter) Status() int { if !r.wroteHeader { return http.StatusOK } + return r.status } diff --git a/schema_test.go b/schema_test.go index 08bb050..e81fd86 100644 --- a/schema_test.go +++ b/schema_test.go @@ -31,6 +31,7 @@ func TestSchema(t *testing.T) { require.NoError(t, err) filename := path.Join(cwd, "gotsrpc.schema.json") + expected, err := os.ReadFile(filename) if !errors.Is(err, os.ErrNotExist) { require.NoError(t, err) diff --git a/tests/aliases/generate_test.go b/tests/aliases/generate_test.go index 5ab5855..3d3be5d 100644 --- a/tests/aliases/generate_test.go +++ b/tests/aliases/generate_test.go @@ -15,6 +15,7 @@ func TestClient(t *testing.T) { cmd := exec.CommandContext(t.Context(), "bun", "test", "./client/client.test.ts") cmd.Stdout = os.Stdout cmd.Stderr = os.Stderr + cmd.Env = append(os.Environ(), "GOTSRPC_SERVER_URL="+s.URL) require.NoError(t, cmd.Run()) } diff --git a/tests/aliases/server/gotsrpcclient_test.go b/tests/aliases/server/gotsrpcclient_test.go index 07c5e47..612f278 100644 --- a/tests/aliases/server/gotsrpcclient_test.go +++ b/tests/aliases/server/gotsrpcclient_test.go @@ -46,6 +46,7 @@ func TestNewDefaultServiceGoTSRPCClient(t *testing.T) { t.Run("TagsValue", func(t *testing.T) { t.Parallel() + v := server.Tags{"go", "typescript", "rpc"} ret, clientErr := c.TagsValue(t.Context(), v) require.NoError(t, clientErr) @@ -54,6 +55,7 @@ func TestNewDefaultServiceGoTSRPCClient(t *testing.T) { t.Run("EntriesValue", func(t *testing.T) { t.Parallel() + e := &server.Entry{ ID: "1", Status: server.StatusActive, Priority: server.PriorityHigh, Rating: 4.5, Tags: server.Tags{"a"}, @@ -68,6 +70,7 @@ func TestNewDefaultServiceGoTSRPCClient(t *testing.T) { t.Run("RegistryValue", func(t *testing.T) { t.Parallel() + v := server.Registry{ "first": {ID: "1", Status: server.StatusActive, Priority: server.PriorityLow, Rating: 1.0}, } @@ -78,6 +81,7 @@ func TestNewDefaultServiceGoTSRPCClient(t *testing.T) { t.Run("IndexValue", func(t *testing.T) { t.Parallel() + v := server.Index{ server.CategoryA: {{ID: "1", Status: server.StatusActive, Priority: server.PriorityLow, Rating: 1.0}}, } @@ -88,6 +92,7 @@ func TestNewDefaultServiceGoTSRPCClient(t *testing.T) { t.Run("LabelMapValue", func(t *testing.T) { t.Parallel() + v := server.LabelMap{"key": "val", "env": "prod"} ret, clientErr := c.LabelMapValue(t.Context(), v) require.NoError(t, clientErr) @@ -96,6 +101,7 @@ func TestNewDefaultServiceGoTSRPCClient(t *testing.T) { t.Run("EntryValue", func(t *testing.T) { t.Parallel() + v := server.Entry{ ID: "1", Status: server.StatusActive, Priority: server.PriorityHigh, Rating: 4.5, Tags: server.Tags{"a", "b"}, @@ -107,6 +113,7 @@ func TestNewDefaultServiceGoTSRPCClient(t *testing.T) { t.Run("DetailValue", func(t *testing.T) { t.Parallel() + v := server.Detail{ Name: "test", Description: "desc", @@ -122,6 +129,7 @@ func TestNewDefaultServiceGoTSRPCClient(t *testing.T) { t.Run("DataRecordValue", func(t *testing.T) { t.Parallel() + note := "some note" v := server.DataRecord{ ID: "1", @@ -153,6 +161,7 @@ func TestNewDefaultServiceGoTSRPCClient(t *testing.T) { t.Run("MapOfEntries", func(t *testing.T) { t.Parallel() + v := map[server.Category][]server.Entry{ server.CategoryA: {{ID: "1", Status: server.StatusActive, Priority: server.PriorityLow, Rating: 1.0}}, } @@ -163,6 +172,7 @@ func TestNewDefaultServiceGoTSRPCClient(t *testing.T) { t.Run("DataRecordNil", func(t *testing.T) { t.Parallel() + v := server.DataRecord{ ID: "1", Title: "minimal", diff --git a/tests/context/generate_test.go b/tests/context/generate_test.go index afa6418..0781e7a 100644 --- a/tests/context/generate_test.go +++ b/tests/context/generate_test.go @@ -15,6 +15,7 @@ func TestClient(t *testing.T) { cmd := exec.CommandContext(t.Context(), "bun", "test", "./client/client.test.ts") cmd.Stdout = os.Stdout cmd.Stderr = os.Stderr + cmd.Env = append(os.Environ(), "GOTSRPC_SERVER_URL="+s.URL) require.NoError(t, cmd.Run()) } diff --git a/tests/errors/generate_test.go b/tests/errors/generate_test.go index 049f90c..4839994 100644 --- a/tests/errors/generate_test.go +++ b/tests/errors/generate_test.go @@ -15,6 +15,7 @@ func TestClient(t *testing.T) { cmd := exec.CommandContext(t.Context(), "bun", "test", "./client/client.test.ts") cmd.Stdout = os.Stdout cmd.Stderr = os.Stderr + cmd.Env = append(os.Environ(), "GOTSRPC_SERVER_URL="+s.URL) require.NoError(t, cmd.Run()) } diff --git a/tests/errors/server/gotsrpcclient_test.go b/tests/errors/server/gotsrpcclient_test.go index e8214cd..66f8773 100644 --- a/tests/errors/server/gotsrpcclient_test.go +++ b/tests/errors/server/gotsrpcclient_test.go @@ -43,6 +43,7 @@ func TestNewDefaultServiceGoTSRPCClient(t *testing.T) { t.Parallel() serviceErr, clientErr := c.Scalar(t.Context()) require.NoError(t, clientErr) + var err *server.ScalarError require.ErrorAs(t, serviceErr, &err) assert.Equal(t, "one", err.String()) diff --git a/tests/errors/server/handler_test.go b/tests/errors/server/handler_test.go index ff6ea1e..e95718d 100644 --- a/tests/errors/server/handler_test.go +++ b/tests/errors/server/handler_test.go @@ -19,12 +19,14 @@ func TestHandler(t *testing.T) { t.Run("Error", func(t *testing.T) { t.Parallel() + err := h.Error(w, r) assert.Error(t, err) }) t.Run("Errors", func(t *testing.T) { t.Parallel() + err1, err2 := h.Errors(w, r) require.Error(t, err1) assert.Error(t, err2) @@ -32,66 +34,77 @@ func TestHandler(t *testing.T) { t.Run("Scalar", func(t *testing.T) { t.Parallel() + ret := h.Scalar(w, r) assert.NotNil(t, ret) }) t.Run("MultiScalar", func(t *testing.T) { t.Parallel() + ret := h.MultiScalar(w, r) assert.NotNil(t, ret) }) t.Run("Struct", func(t *testing.T) { t.Parallel() + ret := h.Struct(w, r) assert.NotNil(t, ret) }) t.Run("StructError", func(t *testing.T) { t.Parallel() + err := h.StructError(w, r) assert.Error(t, err) }) t.Run("TypedError", func(t *testing.T) { t.Parallel() + err := h.TypedError(w, r) assert.Error(t, err) }) t.Run("ScalarError", func(t *testing.T) { t.Parallel() + err := h.ScalarError(w, r) assert.Error(t, err) }) t.Run("CustomError", func(t *testing.T) { t.Parallel() + err := h.CustomError(w, r) assert.Error(t, err) }) t.Run("WrappedError", func(t *testing.T) { t.Parallel() + err := h.WrappedError(w, r) assert.Error(t, err) }) t.Run("TypedWrappedError", func(t *testing.T) { t.Parallel() + err := h.TypedWrappedError(w, r) assert.Error(t, err) }) t.Run("TypedScalarError", func(t *testing.T) { t.Parallel() + err := h.TypedScalarError(w, r) assert.Error(t, err) }) t.Run("TypedCustomError", func(t *testing.T) { t.Parallel() + err := h.TypedCustomError(w, r) assert.Error(t, err) }) diff --git a/tests/nullable/generate_test.go b/tests/nullable/generate_test.go index b981083..9b88719 100644 --- a/tests/nullable/generate_test.go +++ b/tests/nullable/generate_test.go @@ -15,6 +15,7 @@ func TestClient(t *testing.T) { cmd := exec.CommandContext(t.Context(), "bun", "test", "./client/client.test.ts") cmd.Stdout = os.Stdout cmd.Stderr = os.Stderr + cmd.Env = append(os.Environ(), "GOTSRPC_SERVER_URL="+s.URL) require.NoError(t, cmd.Run()) } diff --git a/tests/nullable/server/gorpcclient_test.go b/tests/nullable/server/gorpcclient_test.go index 4ae6f18..a3147a9 100644 --- a/tests/nullable/server/gorpcclient_test.go +++ b/tests/nullable/server/gorpcclient_test.go @@ -26,6 +26,7 @@ func TestNewServiceGoRPCClient(t *testing.T) { t.Run("VariantA", func(t *testing.T) { t.Parallel() + v := server.Base{B1: "hello", D1: server.ACustomTypeOne} ret, clientErr := c.VariantA(v) require.NoError(t, clientErr) @@ -34,6 +35,7 @@ func TestNewServiceGoRPCClient(t *testing.T) { t.Run("VariantB", func(t *testing.T) { t.Parallel() + v := server.BCustomType("test") ret, clientErr := c.VariantB(v) require.NoError(t, clientErr) @@ -42,6 +44,7 @@ func TestNewServiceGoRPCClient(t *testing.T) { t.Run("VariantE", func(t *testing.T) { t.Parallel() + v := &server.Base{B1: "ptr"} ret, clientErr := c.VariantE(v) require.NoError(t, clientErr) @@ -51,6 +54,7 @@ func TestNewServiceGoRPCClient(t *testing.T) { t.Run("VariantH", func(t *testing.T) { t.Parallel() + i1 := server.Base{B1: "one"} i2 := &server.Base{B1: "two"} i3 := []*server.Base{{B1: "three"}} diff --git a/tests/nullable/server/gotsrpcclient_test.go b/tests/nullable/server/gotsrpcclient_test.go index 60101f1..cbde7f9 100644 --- a/tests/nullable/server/gotsrpcclient_test.go +++ b/tests/nullable/server/gotsrpcclient_test.go @@ -18,6 +18,7 @@ func TestNewDefaultServiceGoTSRPCClient(t *testing.T) { t.Run("VariantA", func(t *testing.T) { t.Parallel() + v := server.Base{B1: "hello", D1: server.ACustomTypeOne} ret, clientErr := c.VariantA(t.Context(), v) require.NoError(t, clientErr) @@ -26,6 +27,7 @@ func TestNewDefaultServiceGoTSRPCClient(t *testing.T) { t.Run("VariantB", func(t *testing.T) { t.Parallel() + v := server.BCustomType("test") ret, clientErr := c.VariantB(t.Context(), v) require.NoError(t, clientErr) @@ -34,6 +36,7 @@ func TestNewDefaultServiceGoTSRPCClient(t *testing.T) { t.Run("VariantC", func(t *testing.T) { t.Parallel() + v := server.BCustomTypes{"a", "b"} ret, clientErr := c.VariantC(t.Context(), v) require.NoError(t, clientErr) @@ -42,6 +45,7 @@ func TestNewDefaultServiceGoTSRPCClient(t *testing.T) { t.Run("VariantD", func(t *testing.T) { t.Parallel() + v := server.BCustomTypesMap{"x": "y"} ret, clientErr := c.VariantD(t.Context(), v) require.NoError(t, clientErr) @@ -50,6 +54,7 @@ func TestNewDefaultServiceGoTSRPCClient(t *testing.T) { t.Run("VariantE", func(t *testing.T) { t.Parallel() + v := &server.Base{B1: "ptr"} ret, clientErr := c.VariantE(t.Context(), v) require.NoError(t, clientErr) @@ -66,6 +71,7 @@ func TestNewDefaultServiceGoTSRPCClient(t *testing.T) { t.Run("VariantF", func(t *testing.T) { t.Parallel() + v := []*server.Base{{B1: "one"}, {B1: "two"}} ret, clientErr := c.VariantF(t.Context(), v) require.NoError(t, clientErr) @@ -76,6 +82,7 @@ func TestNewDefaultServiceGoTSRPCClient(t *testing.T) { t.Run("VariantG", func(t *testing.T) { t.Parallel() + v := map[string]*server.Base{"k": {B1: "val"}} ret, clientErr := c.VariantG(t.Context(), v) require.NoError(t, clientErr) @@ -86,6 +93,7 @@ func TestNewDefaultServiceGoTSRPCClient(t *testing.T) { t.Run("VariantH", func(t *testing.T) { t.Parallel() + i1 := server.Base{B1: "one"} i2 := &server.Base{B1: "two"} i3 := []*server.Base{{B1: "three"}} diff --git a/tests/nullable/server/handler_test.go b/tests/nullable/server/handler_test.go index e587f3f..97fcc8e 100644 --- a/tests/nullable/server/handler_test.go +++ b/tests/nullable/server/handler_test.go @@ -14,30 +14,35 @@ func TestHandler(t *testing.T) { t.Run("VariantA", func(t *testing.T) { t.Parallel() + v := server.Base{B1: "hello", D1: server.ACustomTypeOne} assert.Equal(t, v, h.VariantA(t.Context(), v)) }) t.Run("VariantB", func(t *testing.T) { t.Parallel() + v := server.BCustomType("test") assert.Equal(t, v, h.VariantB(t.Context(), v)) }) t.Run("VariantC", func(t *testing.T) { t.Parallel() + v := server.BCustomTypes{"a", "b"} assert.Equal(t, v, h.VariantC(t.Context(), v)) }) t.Run("VariantD", func(t *testing.T) { t.Parallel() + v := server.BCustomTypesMap{"x": "y"} assert.Equal(t, v, h.VariantD(t.Context(), v)) }) t.Run("VariantE", func(t *testing.T) { t.Parallel() + v := &server.Base{B1: "ptr"} assert.Equal(t, v, h.VariantE(t.Context(), v)) }) @@ -49,18 +54,21 @@ func TestHandler(t *testing.T) { t.Run("VariantF", func(t *testing.T) { t.Parallel() + v := []*server.Base{{B1: "one"}, {B1: "two"}} assert.Equal(t, v, h.VariantF(t.Context(), v)) }) t.Run("VariantG", func(t *testing.T) { t.Parallel() + v := map[string]*server.Base{"k": {B1: "val"}} assert.Equal(t, v, h.VariantG(t.Context(), v)) }) t.Run("VariantH", func(t *testing.T) { t.Parallel() + i1 := server.Base{B1: "one"} i2 := &server.Base{B1: "two"} i3 := []*server.Base{{B1: "three"}} diff --git a/tests/time/generate_test.go b/tests/time/generate_test.go index 3b39b8e..3d0534b 100644 --- a/tests/time/generate_test.go +++ b/tests/time/generate_test.go @@ -15,6 +15,7 @@ func TestClient(t *testing.T) { cmd := exec.CommandContext(t.Context(), "bun", "test", "./client/client.test.ts") cmd.Stdout = os.Stdout cmd.Stderr = os.Stderr + cmd.Env = append(os.Environ(), "GOTSRPC_SERVER_URL="+s.URL) require.NoError(t, cmd.Run()) } diff --git a/tests/time/server/gorpcclient_test.go b/tests/time/server/gorpcclient_test.go index 8f56138..45843f1 100644 --- a/tests/time/server/gorpcclient_test.go +++ b/tests/time/server/gorpcclient_test.go @@ -16,10 +16,12 @@ func TestNewServiceGoRPCClient(t *testing.T) { require.NoError(t, l.Close()) s := server.NewServiceGoRPCProxy(l.Addr().String(), &server.Handler{}, nil) + require.NoError(t, s.Start()) defer s.Stop() c := server.NewServiceGoRPCClient(l.Addr().String(), nil) + c.Start() defer c.Stop() diff --git a/tests/types/generate_test.go b/tests/types/generate_test.go index 0bdae29..558dde7 100644 --- a/tests/types/generate_test.go +++ b/tests/types/generate_test.go @@ -15,6 +15,7 @@ func TestClient(t *testing.T) { cmd := exec.CommandContext(t.Context(), "bun", "test", "./client/client.test.ts") cmd.Stdout = os.Stdout cmd.Stderr = os.Stderr + cmd.Env = append(os.Environ(), "GOTSRPC_SERVER_URL="+s.URL) require.NoError(t, cmd.Run()) } diff --git a/tests/types/server/gorpcclient_test.go b/tests/types/server/gorpcclient_test.go index f2bf23b..d273235 100644 --- a/tests/types/server/gorpcclient_test.go +++ b/tests/types/server/gorpcclient_test.go @@ -27,6 +27,7 @@ func TestNewServiceGoRPCClient(t *testing.T) { t.Run("Bool", func(t *testing.T) { t.Parallel() + ret, clientErr := c.Bool(true) require.NoError(t, clientErr) assert.True(t, ret) @@ -34,6 +35,7 @@ func TestNewServiceGoRPCClient(t *testing.T) { t.Run("Int", func(t *testing.T) { t.Parallel() + ret, clientErr := c.Int(42) require.NoError(t, clientErr) assert.Equal(t, 42, ret) @@ -41,6 +43,7 @@ func TestNewServiceGoRPCClient(t *testing.T) { t.Run("Int64", func(t *testing.T) { t.Parallel() + ret, clientErr := c.Int64(int64(9876543210)) require.NoError(t, clientErr) assert.Equal(t, int64(9876543210), ret) @@ -48,6 +51,7 @@ func TestNewServiceGoRPCClient(t *testing.T) { t.Run("Float64", func(t *testing.T) { t.Parallel() + ret, clientErr := c.Float64(3.14159) require.NoError(t, clientErr) assert.InDelta(t, 3.14159, ret, 1e-10) @@ -55,6 +59,7 @@ func TestNewServiceGoRPCClient(t *testing.T) { t.Run("String", func(t *testing.T) { t.Parallel() + ret, clientErr := c.String("hello world") require.NoError(t, clientErr) assert.Equal(t, "hello world", ret) @@ -62,6 +67,7 @@ func TestNewServiceGoRPCClient(t *testing.T) { t.Run("Int8", func(t *testing.T) { t.Parallel() + ret, clientErr := c.Int8(127) require.NoError(t, clientErr) assert.Equal(t, int8(127), ret) @@ -69,6 +75,7 @@ func TestNewServiceGoRPCClient(t *testing.T) { t.Run("Int16", func(t *testing.T) { t.Parallel() + ret, clientErr := c.Int16(32767) require.NoError(t, clientErr) assert.Equal(t, int16(32767), ret) @@ -76,6 +83,7 @@ func TestNewServiceGoRPCClient(t *testing.T) { t.Run("Int32", func(t *testing.T) { t.Parallel() + ret, clientErr := c.Int32(2147483647) require.NoError(t, clientErr) assert.Equal(t, int32(2147483647), ret) @@ -83,6 +91,7 @@ func TestNewServiceGoRPCClient(t *testing.T) { t.Run("Uint", func(t *testing.T) { t.Parallel() + ret, clientErr := c.Uint(42) require.NoError(t, clientErr) assert.Equal(t, uint(42), ret) @@ -90,6 +99,7 @@ func TestNewServiceGoRPCClient(t *testing.T) { t.Run("Uint8", func(t *testing.T) { t.Parallel() + ret, clientErr := c.Uint8(255) require.NoError(t, clientErr) assert.Equal(t, uint8(255), ret) @@ -97,6 +107,7 @@ func TestNewServiceGoRPCClient(t *testing.T) { t.Run("Uint16", func(t *testing.T) { t.Parallel() + ret, clientErr := c.Uint16(65535) require.NoError(t, clientErr) assert.Equal(t, uint16(65535), ret) @@ -104,6 +115,7 @@ func TestNewServiceGoRPCClient(t *testing.T) { t.Run("Uint32", func(t *testing.T) { t.Parallel() + ret, clientErr := c.Uint32(4294967295) require.NoError(t, clientErr) assert.Equal(t, uint32(4294967295), ret) @@ -111,6 +123,7 @@ func TestNewServiceGoRPCClient(t *testing.T) { t.Run("Uint64", func(t *testing.T) { t.Parallel() + ret, clientErr := c.Uint64(9876543210) require.NoError(t, clientErr) assert.Equal(t, uint64(9876543210), ret) @@ -118,6 +131,7 @@ func TestNewServiceGoRPCClient(t *testing.T) { t.Run("Float32", func(t *testing.T) { t.Parallel() + ret, clientErr := c.Float32(3.14) require.NoError(t, clientErr) assert.InDelta(t, float32(3.14), ret, 1e-5) @@ -125,6 +139,7 @@ func TestNewServiceGoRPCClient(t *testing.T) { t.Run("AllScalarsStruct", func(t *testing.T) { t.Parallel() + v := server.AllScalars{ Int8: 127, Int16: 32767, Int32: 2147483647, Uint: 42, Uint8: 255, Uint16: 65535, Uint32: 4294967295, Uint64: 9876543210, @@ -145,6 +160,7 @@ func TestNewServiceGoRPCClient(t *testing.T) { t.Run("StringPtr", func(t *testing.T) { t.Parallel() + v := "test" ret, clientErr := c.StringPtr(&v) require.NoError(t, clientErr) @@ -154,6 +170,7 @@ func TestNewServiceGoRPCClient(t *testing.T) { t.Run("Int8Ptr", func(t *testing.T) { t.Parallel() + v := int8(127) ret, clientErr := c.Int8Ptr(&v) require.NoError(t, clientErr) @@ -163,6 +180,7 @@ func TestNewServiceGoRPCClient(t *testing.T) { t.Run("Int16Ptr", func(t *testing.T) { t.Parallel() + v := int16(32767) ret, clientErr := c.Int16Ptr(&v) require.NoError(t, clientErr) @@ -172,6 +190,7 @@ func TestNewServiceGoRPCClient(t *testing.T) { t.Run("Int32Ptr", func(t *testing.T) { t.Parallel() + v := int32(2147483647) ret, clientErr := c.Int32Ptr(&v) require.NoError(t, clientErr) @@ -181,6 +200,7 @@ func TestNewServiceGoRPCClient(t *testing.T) { t.Run("UintPtr", func(t *testing.T) { t.Parallel() + v := uint(42) ret, clientErr := c.UintPtr(&v) require.NoError(t, clientErr) @@ -190,6 +210,7 @@ func TestNewServiceGoRPCClient(t *testing.T) { t.Run("Uint8Ptr", func(t *testing.T) { t.Parallel() + v := uint8(255) ret, clientErr := c.Uint8Ptr(&v) require.NoError(t, clientErr) @@ -199,6 +220,7 @@ func TestNewServiceGoRPCClient(t *testing.T) { t.Run("Uint16Ptr", func(t *testing.T) { t.Parallel() + v := uint16(65535) ret, clientErr := c.Uint16Ptr(&v) require.NoError(t, clientErr) @@ -208,6 +230,7 @@ func TestNewServiceGoRPCClient(t *testing.T) { t.Run("Uint32Ptr", func(t *testing.T) { t.Parallel() + v := uint32(4294967295) ret, clientErr := c.Uint32Ptr(&v) require.NoError(t, clientErr) @@ -217,6 +240,7 @@ func TestNewServiceGoRPCClient(t *testing.T) { t.Run("Uint64Ptr", func(t *testing.T) { t.Parallel() + v := uint64(9876543210) ret, clientErr := c.Uint64Ptr(&v) require.NoError(t, clientErr) @@ -226,6 +250,7 @@ func TestNewServiceGoRPCClient(t *testing.T) { t.Run("Float32Ptr", func(t *testing.T) { t.Parallel() + v := float32(3.14) ret, clientErr := c.Float32Ptr(&v) require.NoError(t, clientErr) @@ -235,6 +260,7 @@ func TestNewServiceGoRPCClient(t *testing.T) { t.Run("AllScalarPointersStruct", func(t *testing.T) { t.Parallel() + i8, i16, i32 := int8(127), int16(32767), int32(2147483647) u, u8, u16, u32, u64 := uint(42), uint8(255), uint16(65535), uint32(4294967295), uint64(9876543210) f32 := float32(3.14) @@ -267,6 +293,7 @@ func TestNewServiceGoRPCClient(t *testing.T) { t.Run("SimpleStruct", func(t *testing.T) { t.Parallel() + v := common.Simple{Bool: true, Int: 42, Int64: 100, Float64: 2.718, String: "test"} ret, clientErr := c.SimpleStruct(v) require.NoError(t, clientErr) @@ -275,6 +302,7 @@ func TestNewServiceGoRPCClient(t *testing.T) { t.Run("NestedStruct", func(t *testing.T) { t.Parallel() + v := common.Nested{ Name: "parent", Child: common.Simple{Bool: true, Int: 1, Int64: 2, Float64: 3.0, String: "child"}, @@ -286,6 +314,7 @@ func TestNewServiceGoRPCClient(t *testing.T) { t.Run("InlinedStruct", func(t *testing.T) { t.Parallel() + v := server.Inlined{ Simple: common.Simple{Bool: true, Int: 1, Int64: 2, Float64: 3.0, String: "child"}, Name: "parent", @@ -297,6 +326,7 @@ func TestNewServiceGoRPCClient(t *testing.T) { t.Run("StringSlice", func(t *testing.T) { t.Parallel() + v := []string{"a", "b", "c"} ret, clientErr := c.StringSlice(v) require.NoError(t, clientErr) @@ -305,6 +335,7 @@ func TestNewServiceGoRPCClient(t *testing.T) { t.Run("Int8Slice", func(t *testing.T) { t.Parallel() + v := []int8{-128, 0, 127} ret, clientErr := c.Int8Slice(v) require.NoError(t, clientErr) @@ -313,6 +344,7 @@ func TestNewServiceGoRPCClient(t *testing.T) { t.Run("Int16Slice", func(t *testing.T) { t.Parallel() + v := []int16{-32768, 0, 32767} ret, clientErr := c.Int16Slice(v) require.NoError(t, clientErr) @@ -321,6 +353,7 @@ func TestNewServiceGoRPCClient(t *testing.T) { t.Run("Int32Slice", func(t *testing.T) { t.Parallel() + v := []int32{-2147483648, 0, 2147483647} ret, clientErr := c.Int32Slice(v) require.NoError(t, clientErr) @@ -329,6 +362,7 @@ func TestNewServiceGoRPCClient(t *testing.T) { t.Run("UintSlice", func(t *testing.T) { t.Parallel() + v := []uint{0, 42, 100} ret, clientErr := c.UintSlice(v) require.NoError(t, clientErr) @@ -337,6 +371,7 @@ func TestNewServiceGoRPCClient(t *testing.T) { t.Run("Uint16Slice", func(t *testing.T) { t.Parallel() + v := []uint16{0, 1000, 65535} ret, clientErr := c.Uint16Slice(v) require.NoError(t, clientErr) @@ -345,6 +380,7 @@ func TestNewServiceGoRPCClient(t *testing.T) { t.Run("Uint32Slice", func(t *testing.T) { t.Parallel() + v := []uint32{0, 1000, 4294967295} ret, clientErr := c.Uint32Slice(v) require.NoError(t, clientErr) @@ -353,6 +389,7 @@ func TestNewServiceGoRPCClient(t *testing.T) { t.Run("Uint64Slice", func(t *testing.T) { t.Parallel() + v := []uint64{0, 1000, 9876543210} ret, clientErr := c.Uint64Slice(v) require.NoError(t, clientErr) @@ -361,6 +398,7 @@ func TestNewServiceGoRPCClient(t *testing.T) { t.Run("Float32Slice", func(t *testing.T) { t.Parallel() + v := []float32{1.1, 2.2, 3.3} ret, clientErr := c.Float32Slice(v) require.NoError(t, clientErr) @@ -372,6 +410,7 @@ func TestNewServiceGoRPCClient(t *testing.T) { t.Run("AllScalarSlicesStruct", func(t *testing.T) { t.Parallel() + v := server.AllScalarSlices{ Int8s: []int8{-1, 0, 1}, Int16s: []int16{-1, 0, 1}, Int32s: []int32{-1, 0, 1}, Uints: []uint{0, 1, 2}, Uint16s: []uint16{0, 1, 2}, @@ -394,6 +433,7 @@ func TestNewServiceGoRPCClient(t *testing.T) { t.Run("StringStringMap", func(t *testing.T) { t.Parallel() + v := map[string]string{"a": "1", "b": "2"} ret, clientErr := c.StringStringMap(v) require.NoError(t, clientErr) @@ -402,6 +442,7 @@ func TestNewServiceGoRPCClient(t *testing.T) { t.Run("StringInt8Map", func(t *testing.T) { t.Parallel() + v := map[string]int8{"a": 1, "b": -1} ret, clientErr := c.StringInt8Map(v) require.NoError(t, clientErr) @@ -410,6 +451,7 @@ func TestNewServiceGoRPCClient(t *testing.T) { t.Run("StringInt16Map", func(t *testing.T) { t.Parallel() + v := map[string]int16{"a": 100, "b": -100} ret, clientErr := c.StringInt16Map(v) require.NoError(t, clientErr) @@ -418,6 +460,7 @@ func TestNewServiceGoRPCClient(t *testing.T) { t.Run("StringInt32Map", func(t *testing.T) { t.Parallel() + v := map[string]int32{"a": 100000, "b": -100000} ret, clientErr := c.StringInt32Map(v) require.NoError(t, clientErr) @@ -426,6 +469,7 @@ func TestNewServiceGoRPCClient(t *testing.T) { t.Run("StringUintMap", func(t *testing.T) { t.Parallel() + v := map[string]uint{"a": 0, "b": 42} ret, clientErr := c.StringUintMap(v) require.NoError(t, clientErr) @@ -434,6 +478,7 @@ func TestNewServiceGoRPCClient(t *testing.T) { t.Run("StringUint8Map", func(t *testing.T) { t.Parallel() + v := map[string]uint8{"a": 0, "b": 255} ret, clientErr := c.StringUint8Map(v) require.NoError(t, clientErr) @@ -442,6 +487,7 @@ func TestNewServiceGoRPCClient(t *testing.T) { t.Run("StringUint16Map", func(t *testing.T) { t.Parallel() + v := map[string]uint16{"a": 0, "b": 65535} ret, clientErr := c.StringUint16Map(v) require.NoError(t, clientErr) @@ -450,6 +496,7 @@ func TestNewServiceGoRPCClient(t *testing.T) { t.Run("StringUint32Map", func(t *testing.T) { t.Parallel() + v := map[string]uint32{"a": 0, "b": 4294967295} ret, clientErr := c.StringUint32Map(v) require.NoError(t, clientErr) @@ -458,6 +505,7 @@ func TestNewServiceGoRPCClient(t *testing.T) { t.Run("StringUint64Map", func(t *testing.T) { t.Parallel() + v := map[string]uint64{"a": 0, "b": 9876543210} ret, clientErr := c.StringUint64Map(v) require.NoError(t, clientErr) @@ -466,6 +514,7 @@ func TestNewServiceGoRPCClient(t *testing.T) { t.Run("StringFloat32Map", func(t *testing.T) { t.Parallel() + v := map[string]float32{"a": 1.1, "b": 2.2} ret, clientErr := c.StringFloat32Map(v) require.NoError(t, clientErr) @@ -477,6 +526,7 @@ func TestNewServiceGoRPCClient(t *testing.T) { t.Run("AllScalarMapsStruct", func(t *testing.T) { t.Parallel() + v := server.AllScalarMaps{ Int8Map: map[string]int8{"x": 1}, Int16Map: map[string]int16{"x": 1}, Int32Map: map[string]int32{"x": 1}, UintMap: map[string]uint{"x": 1}, @@ -500,6 +550,7 @@ func TestNewServiceGoRPCClient(t *testing.T) { t.Run("MapOfMaps", func(t *testing.T) { t.Parallel() + v := map[string]map[string]string{"outer": {"inner": "val"}} ret, clientErr := c.MapOfMaps(v) require.NoError(t, clientErr) @@ -508,6 +559,7 @@ func TestNewServiceGoRPCClient(t *testing.T) { t.Run("MultiArgs", func(t *testing.T) { t.Parallel() + retA, retB, retC, clientErr := c.MultiArgs("hello", int64(42), true) require.NoError(t, clientErr) assert.Equal(t, "hello", retA) @@ -517,6 +569,7 @@ func TestNewServiceGoRPCClient(t *testing.T) { t.Run("Empty", func(t *testing.T) { t.Parallel() + ret, clientErr := c.Empty() require.NoError(t, clientErr) assert.True(t, ret) @@ -524,6 +577,7 @@ func TestNewServiceGoRPCClient(t *testing.T) { t.Run("ByteSlice", func(t *testing.T) { t.Parallel() + v := []byte("hello world") ret, clientErr := c.ByteSlice(v) require.NoError(t, clientErr) @@ -532,6 +586,7 @@ func TestNewServiceGoRPCClient(t *testing.T) { t.Run("ObjectID", func(t *testing.T) { t.Parallel() + var v server.ObjectID copy(v[:], "hello123456") ret, clientErr := c.ObjectID(v) diff --git a/tests/types/server/gotsrpcclient_test.go b/tests/types/server/gotsrpcclient_test.go index d6d24a7..f9f3d7d 100644 --- a/tests/types/server/gotsrpcclient_test.go +++ b/tests/types/server/gotsrpcclient_test.go @@ -117,6 +117,7 @@ func TestNewDefaultServiceGoTSRPCClient(t *testing.T) { t.Run("AllScalarsStruct", func(t *testing.T) { t.Parallel() + v := server.AllScalars{ Int8: 127, Int16: 32767, Int32: 2147483647, Uint: 42, Uint8: 255, Uint16: 65535, Uint32: 4294967295, Uint64: 9876543210, @@ -137,6 +138,7 @@ func TestNewDefaultServiceGoTSRPCClient(t *testing.T) { t.Run("StringPtr", func(t *testing.T) { t.Parallel() + v := "test" ret, clientErr := c.StringPtr(t.Context(), &v) require.NoError(t, clientErr) @@ -146,6 +148,7 @@ func TestNewDefaultServiceGoTSRPCClient(t *testing.T) { t.Run("Int64Ptr", func(t *testing.T) { t.Parallel() + v := int64(42) ret, clientErr := c.Int64Ptr(t.Context(), &v) require.NoError(t, clientErr) @@ -155,6 +158,7 @@ func TestNewDefaultServiceGoTSRPCClient(t *testing.T) { t.Run("BoolPtr", func(t *testing.T) { t.Parallel() + v := true ret, clientErr := c.BoolPtr(t.Context(), &v) require.NoError(t, clientErr) @@ -164,6 +168,7 @@ func TestNewDefaultServiceGoTSRPCClient(t *testing.T) { t.Run("Int8Ptr", func(t *testing.T) { t.Parallel() + v := int8(127) ret, clientErr := c.Int8Ptr(t.Context(), &v) require.NoError(t, clientErr) @@ -173,6 +178,7 @@ func TestNewDefaultServiceGoTSRPCClient(t *testing.T) { t.Run("Int16Ptr", func(t *testing.T) { t.Parallel() + v := int16(32767) ret, clientErr := c.Int16Ptr(t.Context(), &v) require.NoError(t, clientErr) @@ -182,6 +188,7 @@ func TestNewDefaultServiceGoTSRPCClient(t *testing.T) { t.Run("Int32Ptr", func(t *testing.T) { t.Parallel() + v := int32(2147483647) ret, clientErr := c.Int32Ptr(t.Context(), &v) require.NoError(t, clientErr) @@ -191,6 +198,7 @@ func TestNewDefaultServiceGoTSRPCClient(t *testing.T) { t.Run("UintPtr", func(t *testing.T) { t.Parallel() + v := uint(42) ret, clientErr := c.UintPtr(t.Context(), &v) require.NoError(t, clientErr) @@ -200,6 +208,7 @@ func TestNewDefaultServiceGoTSRPCClient(t *testing.T) { t.Run("Uint8Ptr", func(t *testing.T) { t.Parallel() + v := uint8(255) ret, clientErr := c.Uint8Ptr(t.Context(), &v) require.NoError(t, clientErr) @@ -209,6 +218,7 @@ func TestNewDefaultServiceGoTSRPCClient(t *testing.T) { t.Run("Uint16Ptr", func(t *testing.T) { t.Parallel() + v := uint16(65535) ret, clientErr := c.Uint16Ptr(t.Context(), &v) require.NoError(t, clientErr) @@ -218,6 +228,7 @@ func TestNewDefaultServiceGoTSRPCClient(t *testing.T) { t.Run("Uint32Ptr", func(t *testing.T) { t.Parallel() + v := uint32(4294967295) ret, clientErr := c.Uint32Ptr(t.Context(), &v) require.NoError(t, clientErr) @@ -227,6 +238,7 @@ func TestNewDefaultServiceGoTSRPCClient(t *testing.T) { t.Run("Uint64Ptr", func(t *testing.T) { t.Parallel() + v := uint64(9876543210) ret, clientErr := c.Uint64Ptr(t.Context(), &v) require.NoError(t, clientErr) @@ -236,6 +248,7 @@ func TestNewDefaultServiceGoTSRPCClient(t *testing.T) { t.Run("Float32Ptr", func(t *testing.T) { t.Parallel() + v := float32(3.14) ret, clientErr := c.Float32Ptr(t.Context(), &v) require.NoError(t, clientErr) @@ -245,6 +258,7 @@ func TestNewDefaultServiceGoTSRPCClient(t *testing.T) { t.Run("AllScalarPointersStruct", func(t *testing.T) { t.Parallel() + i8, i16, i32 := int8(127), int16(32767), int32(2147483647) u, u8, u16, u32, u64 := uint(42), uint8(255), uint16(65535), uint32(4294967295), uint64(9876543210) f32 := float32(3.14) @@ -277,6 +291,7 @@ func TestNewDefaultServiceGoTSRPCClient(t *testing.T) { t.Run("SimpleStruct", func(t *testing.T) { t.Parallel() + v := common.Simple{ Bool: true, Int: 42, Int64: 100, Float64: 2.718, String: "test", } @@ -287,6 +302,7 @@ func TestNewDefaultServiceGoTSRPCClient(t *testing.T) { t.Run("NestedStruct", func(t *testing.T) { t.Parallel() + v := common.Nested{ Name: "parent", Child: common.Simple{Bool: true, Int: 1, Int64: 2, Float64: 3.0, String: "child"}, @@ -298,6 +314,7 @@ func TestNewDefaultServiceGoTSRPCClient(t *testing.T) { t.Run("InlinedStruct", func(t *testing.T) { t.Parallel() + v := server.Inlined{ Simple: common.Simple{Bool: true, Int: 1, Int64: 2, Float64: 3.0, String: "child"}, Name: "parent", @@ -309,6 +326,7 @@ func TestNewDefaultServiceGoTSRPCClient(t *testing.T) { t.Run("StructWithPointers", func(t *testing.T) { t.Parallel() + str := "hello" i := int64(42) b := true @@ -330,6 +348,7 @@ func TestNewDefaultServiceGoTSRPCClient(t *testing.T) { t.Run("StructWithCollections", func(t *testing.T) { t.Parallel() + v := server.WithCollections{ Strings: []string{"a", "b"}, Int64s: []int64{1, 2, 3}, @@ -352,6 +371,7 @@ func TestNewDefaultServiceGoTSRPCClient(t *testing.T) { t.Run("StringSlice", func(t *testing.T) { t.Parallel() + v := []string{"a", "b", "c"} ret, clientErr := c.StringSlice(t.Context(), v) require.NoError(t, clientErr) @@ -360,6 +380,7 @@ func TestNewDefaultServiceGoTSRPCClient(t *testing.T) { t.Run("Int64Slice", func(t *testing.T) { t.Parallel() + v := []int64{10, 20, 30} ret, clientErr := c.Int64Slice(t.Context(), v) require.NoError(t, clientErr) @@ -368,6 +389,7 @@ func TestNewDefaultServiceGoTSRPCClient(t *testing.T) { t.Run("SimpleSlice", func(t *testing.T) { t.Parallel() + v := []common.Simple{ {Bool: true, Int: 1, Int64: 2, Float64: 3.0, String: "one"}, {Bool: false, Int: 4, Int64: 5, Float64: 6.0, String: "two"}, @@ -379,6 +401,7 @@ func TestNewDefaultServiceGoTSRPCClient(t *testing.T) { t.Run("SimplePtrSlice", func(t *testing.T) { t.Parallel() + s1 := common.Simple{Bool: true, Int: 1, Int64: 2, Float64: 3.0, String: "one"} s2 := common.Simple{Bool: false, Int: 4, Int64: 5, Float64: 6.0, String: "two"} v := []*common.Simple{&s1, &s2} @@ -393,6 +416,7 @@ func TestNewDefaultServiceGoTSRPCClient(t *testing.T) { t.Run("StringSlice2D", func(t *testing.T) { t.Parallel() + v := [][]string{{"a", "b"}, {"c", "d"}} ret, clientErr := c.StringSlice2D(t.Context(), v) require.NoError(t, clientErr) @@ -401,6 +425,7 @@ func TestNewDefaultServiceGoTSRPCClient(t *testing.T) { t.Run("Int8Slice", func(t *testing.T) { t.Parallel() + v := []int8{-128, 0, 127} ret, clientErr := c.Int8Slice(t.Context(), v) require.NoError(t, clientErr) @@ -409,6 +434,7 @@ func TestNewDefaultServiceGoTSRPCClient(t *testing.T) { t.Run("Int16Slice", func(t *testing.T) { t.Parallel() + v := []int16{-32768, 0, 32767} ret, clientErr := c.Int16Slice(t.Context(), v) require.NoError(t, clientErr) @@ -417,6 +443,7 @@ func TestNewDefaultServiceGoTSRPCClient(t *testing.T) { t.Run("Int32Slice", func(t *testing.T) { t.Parallel() + v := []int32{-2147483648, 0, 2147483647} ret, clientErr := c.Int32Slice(t.Context(), v) require.NoError(t, clientErr) @@ -425,6 +452,7 @@ func TestNewDefaultServiceGoTSRPCClient(t *testing.T) { t.Run("UintSlice", func(t *testing.T) { t.Parallel() + v := []uint{0, 42, 100} ret, clientErr := c.UintSlice(t.Context(), v) require.NoError(t, clientErr) @@ -433,6 +461,7 @@ func TestNewDefaultServiceGoTSRPCClient(t *testing.T) { t.Run("Uint16Slice", func(t *testing.T) { t.Parallel() + v := []uint16{0, 1000, 65535} ret, clientErr := c.Uint16Slice(t.Context(), v) require.NoError(t, clientErr) @@ -441,6 +470,7 @@ func TestNewDefaultServiceGoTSRPCClient(t *testing.T) { t.Run("Uint32Slice", func(t *testing.T) { t.Parallel() + v := []uint32{0, 1000, 4294967295} ret, clientErr := c.Uint32Slice(t.Context(), v) require.NoError(t, clientErr) @@ -449,6 +479,7 @@ func TestNewDefaultServiceGoTSRPCClient(t *testing.T) { t.Run("Uint64Slice", func(t *testing.T) { t.Parallel() + v := []uint64{0, 1000, 9876543210} ret, clientErr := c.Uint64Slice(t.Context(), v) require.NoError(t, clientErr) @@ -457,6 +488,7 @@ func TestNewDefaultServiceGoTSRPCClient(t *testing.T) { t.Run("Float32Slice", func(t *testing.T) { t.Parallel() + v := []float32{1.1, 2.2, 3.3} ret, clientErr := c.Float32Slice(t.Context(), v) require.NoError(t, clientErr) @@ -468,6 +500,7 @@ func TestNewDefaultServiceGoTSRPCClient(t *testing.T) { t.Run("AllScalarSlicesStruct", func(t *testing.T) { t.Parallel() + v := server.AllScalarSlices{ Int8s: []int8{-1, 0, 1}, Int16s: []int16{-1, 0, 1}, Int32s: []int32{-1, 0, 1}, Uints: []uint{0, 1, 2}, Uint16s: []uint16{0, 1, 2}, @@ -490,6 +523,7 @@ func TestNewDefaultServiceGoTSRPCClient(t *testing.T) { t.Run("StringStringMap", func(t *testing.T) { t.Parallel() + v := map[string]string{"a": "1", "b": "2"} ret, clientErr := c.StringStringMap(t.Context(), v) require.NoError(t, clientErr) @@ -498,6 +532,7 @@ func TestNewDefaultServiceGoTSRPCClient(t *testing.T) { t.Run("StringInt64Map", func(t *testing.T) { t.Parallel() + v := map[string]int64{"x": 10, "y": 20} ret, clientErr := c.StringInt64Map(t.Context(), v) require.NoError(t, clientErr) @@ -506,6 +541,7 @@ func TestNewDefaultServiceGoTSRPCClient(t *testing.T) { t.Run("StringSimpleMap", func(t *testing.T) { t.Parallel() + v := map[string]common.Simple{ "one": {Bool: true, Int: 1, Int64: 2, Float64: 3.0, String: "one"}, } @@ -516,6 +552,7 @@ func TestNewDefaultServiceGoTSRPCClient(t *testing.T) { t.Run("StringSimplePtrMap", func(t *testing.T) { t.Parallel() + s := common.Simple{Bool: true, Int: 1, Int64: 2, Float64: 3.0, String: "val"} v := map[string]*common.Simple{"k": &s} ret, clientErr := c.StringSimplePtrMap(t.Context(), v) @@ -527,6 +564,7 @@ func TestNewDefaultServiceGoTSRPCClient(t *testing.T) { t.Run("StringStringSliceMap", func(t *testing.T) { t.Parallel() + v := map[string][]string{"colors": {"red", "blue"}, "sizes": {"s", "m"}} ret, clientErr := c.StringStringSliceMap(t.Context(), v) require.NoError(t, clientErr) @@ -535,6 +573,7 @@ func TestNewDefaultServiceGoTSRPCClient(t *testing.T) { t.Run("StringInt8Map", func(t *testing.T) { t.Parallel() + v := map[string]int8{"a": 1, "b": -1} ret, clientErr := c.StringInt8Map(t.Context(), v) require.NoError(t, clientErr) @@ -543,6 +582,7 @@ func TestNewDefaultServiceGoTSRPCClient(t *testing.T) { t.Run("StringInt16Map", func(t *testing.T) { t.Parallel() + v := map[string]int16{"a": 100, "b": -100} ret, clientErr := c.StringInt16Map(t.Context(), v) require.NoError(t, clientErr) @@ -551,6 +591,7 @@ func TestNewDefaultServiceGoTSRPCClient(t *testing.T) { t.Run("StringInt32Map", func(t *testing.T) { t.Parallel() + v := map[string]int32{"a": 100000, "b": -100000} ret, clientErr := c.StringInt32Map(t.Context(), v) require.NoError(t, clientErr) @@ -559,6 +600,7 @@ func TestNewDefaultServiceGoTSRPCClient(t *testing.T) { t.Run("StringUintMap", func(t *testing.T) { t.Parallel() + v := map[string]uint{"a": 0, "b": 42} ret, clientErr := c.StringUintMap(t.Context(), v) require.NoError(t, clientErr) @@ -567,6 +609,7 @@ func TestNewDefaultServiceGoTSRPCClient(t *testing.T) { t.Run("StringUint8Map", func(t *testing.T) { t.Parallel() + v := map[string]uint8{"a": 0, "b": 255} ret, clientErr := c.StringUint8Map(t.Context(), v) require.NoError(t, clientErr) @@ -575,6 +618,7 @@ func TestNewDefaultServiceGoTSRPCClient(t *testing.T) { t.Run("StringUint16Map", func(t *testing.T) { t.Parallel() + v := map[string]uint16{"a": 0, "b": 65535} ret, clientErr := c.StringUint16Map(t.Context(), v) require.NoError(t, clientErr) @@ -583,6 +627,7 @@ func TestNewDefaultServiceGoTSRPCClient(t *testing.T) { t.Run("StringUint32Map", func(t *testing.T) { t.Parallel() + v := map[string]uint32{"a": 0, "b": 4294967295} ret, clientErr := c.StringUint32Map(t.Context(), v) require.NoError(t, clientErr) @@ -591,6 +636,7 @@ func TestNewDefaultServiceGoTSRPCClient(t *testing.T) { t.Run("StringUint64Map", func(t *testing.T) { t.Parallel() + v := map[string]uint64{"a": 0, "b": 9876543210} ret, clientErr := c.StringUint64Map(t.Context(), v) require.NoError(t, clientErr) @@ -599,6 +645,7 @@ func TestNewDefaultServiceGoTSRPCClient(t *testing.T) { t.Run("StringFloat32Map", func(t *testing.T) { t.Parallel() + v := map[string]float32{"a": 1.1, "b": 2.2} ret, clientErr := c.StringFloat32Map(t.Context(), v) require.NoError(t, clientErr) @@ -610,6 +657,7 @@ func TestNewDefaultServiceGoTSRPCClient(t *testing.T) { t.Run("AllScalarMapsStruct", func(t *testing.T) { t.Parallel() + v := server.AllScalarMaps{ Int8Map: map[string]int8{"x": 1}, Int16Map: map[string]int16{"x": 1}, Int32Map: map[string]int32{"x": 1}, UintMap: map[string]uint{"x": 1}, @@ -633,6 +681,7 @@ func TestNewDefaultServiceGoTSRPCClient(t *testing.T) { t.Run("MapOfMaps", func(t *testing.T) { t.Parallel() + v := map[string]map[string]string{"outer": {"inner": "val"}} ret, clientErr := c.MapOfMaps(t.Context(), v) require.NoError(t, clientErr) @@ -641,6 +690,7 @@ func TestNewDefaultServiceGoTSRPCClient(t *testing.T) { t.Run("MapOfSimpleSlice", func(t *testing.T) { t.Parallel() + v := map[string][]common.Simple{ "group": {{Bool: true, Int: 1, Int64: 2, Float64: 3.0, String: "item"}}, } @@ -651,6 +701,7 @@ func TestNewDefaultServiceGoTSRPCClient(t *testing.T) { t.Run("SliceOfMaps", func(t *testing.T) { t.Parallel() + v := []map[string]string{{"a": "1"}, {"b": "2"}} ret, clientErr := c.SliceOfMaps(t.Context(), v) require.NoError(t, clientErr) @@ -668,6 +719,7 @@ func TestNewDefaultServiceGoTSRPCClient(t *testing.T) { t.Run("MixedArgs", func(t *testing.T) { t.Parallel() + s := common.Simple{Bool: true, Int: 1, Int64: 2, Float64: 3.0, String: "mix"} items := []string{"a", "b"} m := map[string]int64{"x": 10} @@ -687,6 +739,7 @@ func TestNewDefaultServiceGoTSRPCClient(t *testing.T) { t.Run("ByteSlice", func(t *testing.T) { t.Parallel() + v := []byte("hello world") ret, clientErr := c.ByteSlice(t.Context(), v) require.NoError(t, clientErr) @@ -695,6 +748,7 @@ func TestNewDefaultServiceGoTSRPCClient(t *testing.T) { t.Run("ObjectID", func(t *testing.T) { t.Parallel() + var v server.ObjectID copy(v[:], "hello123456") ret, clientErr := c.ObjectID(t.Context(), v) diff --git a/tests/union/generate_test.go b/tests/union/generate_test.go index 9a2da96..6a1dd54 100644 --- a/tests/union/generate_test.go +++ b/tests/union/generate_test.go @@ -15,6 +15,7 @@ func TestClient(t *testing.T) { cmd := exec.CommandContext(t.Context(), "bun", "test", "./client/client.test.ts") cmd.Stdout = os.Stdout cmd.Stderr = os.Stderr + cmd.Env = append(os.Environ(), "GOTSRPC_SERVER_URL="+s.URL) require.NoError(t, cmd.Run()) } diff --git a/transporthandle.go b/transporthandle.go index 54d4f6a..faab4be 100644 --- a/transporthandle.go +++ b/transporthandle.go @@ -24,6 +24,7 @@ func (ch *transportHandle) getEncoder(w io.Writer) *codec.Encoder { enc.Reset(w) return enc } + return codec.NewEncoder(w, ch.handle) } @@ -36,6 +37,7 @@ func (ch *transportHandle) getDecoder(r io.Reader) *codec.Decoder { dec.Reset(r) return dec } + return codec.NewDecoder(r, ch.handle) } @@ -55,6 +57,7 @@ var ( func registerTransportHandle(encoding ClientEncoding, h *transportHandle) { handlesByEncoding[encoding] = h + handlesByContentType[h.contentType] = h if defaultTransportHandle == nil { defaultTransportHandle = h @@ -76,6 +79,7 @@ func newErrorEncodeHook() func(*[]any, []int) error { (*resp)[i] = NewError(e) } } + return nil } } @@ -85,12 +89,16 @@ func newErrorDecodeHook() func([]any, []int) ([]any, error) { if len(errorIndices) == 0 { return reply, nil } + ret := make([]any, len(reply)) copy(ret, reply) + for _, i := range errorIndices { var e *Error + ret[i] = e } + return ret, nil } } @@ -106,6 +114,7 @@ func newErrorAfterDecodeHook() func(*[]any, []any, []int) error { } } } + return nil } } @@ -118,6 +127,7 @@ func getHandleForEncoding(encoding ClientEncoding) *transportHandle { if h, ok := handlesByEncoding[encoding]; ok { return h } + return defaultTransportHandle } @@ -125,5 +135,6 @@ func getHandlerForContentType(contentType string) *transportHandle { if h, ok := handlesByContentType[contentType]; ok { return h } + return defaultTransportHandle } diff --git a/transportjson.go b/transportjson.go index 56d0016..d2f6ffa 100644 --- a/transportjson.go +++ b/transportjson.go @@ -31,5 +31,6 @@ func SetJSONExt(rt interface{}, tag uint64, ext codec.InterfaceExt) error { if value, ok := jsonHandle.handle.(*codec.JsonHandle); ok { return value.SetInterfaceExt(reflect.TypeOf(rt), tag, ext) } + return errors.New("invalid handle type") } diff --git a/transportmsgpack.go b/transportmsgpack.go index dab9eb9..f5a91d5 100644 --- a/transportmsgpack.go +++ b/transportmsgpack.go @@ -40,5 +40,6 @@ func SetMSGPackExt(rt interface{}, tag uint64, ext codec.BytesExt) error { if value, ok := msgpackHandle.handle.(*codec.MsgpackHandle); ok { return value.SetBytesExt(reflect.TypeOf(rt), tag, ext) } + return errors.New("invalid handle type") } diff --git a/unionext.go b/unionext.go index 3f39bb5..a5e8330 100644 --- a/unionext.go +++ b/unionext.go @@ -15,6 +15,7 @@ func RegisterUnionExt(v ...interface{}) error { return err } } + return nil } @@ -29,11 +30,13 @@ func (x *UnionExt) ConvertExt(v interface{}) interface{} { if val.Kind() == reflect.Ptr { val = val.Elem() } + for i := 0; i < val.NumField(); i++ { if field := val.Field(i); !field.IsZero() { return field.Interface() } } + return nil }