Skip to content

Commit 9e520d1

Browse files
committed
Emit Rust type params for generic interface functions
1 parent a39ec54 commit 9e520d1

3 files changed

Lines changed: 112 additions & 0 deletions

File tree

‎go/decl.go‎

Lines changed: 41 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -982,6 +982,46 @@ func isOsArgsSelector(sel *ast.SelectorExpr) bool {
982982
return resolveStdlibPackageName(ident.Name) == "os"
983983
}
984984

985+
func writeFunctionTypeParams(out *strings.Builder, fnType *ast.FuncType) {
986+
if fnType == nil || fnType.TypeParams == nil || len(fnType.TypeParams.List) == 0 {
987+
return
988+
}
989+
var params []string
990+
for _, field := range fnType.TypeParams.List {
991+
for _, name := range field.Names {
992+
params = append(params, rustFunctionTypeParam(name))
993+
}
994+
}
995+
if len(params) == 0 {
996+
return
997+
}
998+
out.WriteString("<")
999+
out.WriteString(strings.Join(params, ", "))
1000+
out.WriteString(">")
1001+
}
1002+
1003+
func rustFunctionTypeParam(name *ast.Ident) string {
1004+
rustName := RustTypeNameForUse(name.Name)
1005+
typeInfo := GetTypeInfo()
1006+
if typeInfo == nil || typeInfo.info == nil {
1007+
return rustName
1008+
}
1009+
obj, ok := typeInfo.info.Defs[name].(*types.TypeName)
1010+
if !ok {
1011+
return rustName
1012+
}
1013+
traitName, ok := goTypeParamTraitConstraintName(obj.Type())
1014+
if !ok {
1015+
return rustName
1016+
}
1017+
bounds := []string{traitName, "Clone"}
1018+
if NeedsConcurrentWrapper() {
1019+
bounds = append(bounds, "Send", "Sync")
1020+
}
1021+
bounds = append(bounds, "'static")
1022+
return rustName + ": " + strings.Join(bounds, " + ")
1023+
}
1024+
9851025
func functionUsesOsArgs(fn *ast.FuncDecl) bool {
9861026
if fn.Body == nil {
9871027
return false
@@ -1530,6 +1570,7 @@ func TranspileFunction(out *strings.Builder, fn *ast.FuncDecl, fileSet *token.Fi
15301570
}
15311571
out.WriteString("fn ")
15321572
out.WriteString(rustFunctionName(fn))
1573+
writeFunctionTypeParams(out, fn.Type)
15331574
out.WriteString("(")
15341575

15351576
// Parameters

‎go/decl_test.go‎

Lines changed: 40 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -2,6 +2,7 @@ package main
22

33
import (
44
"go/ast"
5+
"go/parser"
56
"go/token"
67
"strings"
78
"testing"
@@ -30,6 +31,45 @@ func TestTranspileFunctionWithoutBodyDoesNotPanic(t *testing.T) {
3031
}
3132
}
3233

34+
func TestTranspileGenericInterfaceConstrainedFunctionEmitsRustTypeParam(t *testing.T) {
35+
fset := token.NewFileSet()
36+
file, err := parser.ParseFile(fset, "main.go", `package main
37+
38+
type Node interface {
39+
Pos() int
40+
}
41+
42+
func Use(node Node) {}
43+
44+
func VisitAll[N Node](list []N) {
45+
for _, node := range list {
46+
Use(node)
47+
}
48+
}
49+
`, 0)
50+
if err != nil {
51+
t.Fatalf("ParseFile() error = %v", err)
52+
}
53+
typeInfo, err := NewTypeInfo([]*ast.File{file}, fset)
54+
if err != nil {
55+
t.Fatalf("NewTypeInfo() error = %v", err)
56+
}
57+
SetTypeInfo(typeInfo)
58+
defer SetTypeInfo(nil)
59+
60+
rust, _, _ := Transpile(file, fset, typeInfo)
61+
62+
if !strings.Contains(rust, "pub fn visit_all<N: Node + Clone") {
63+
t.Fatalf("generic interface-constrained function should emit a Rust type parameter bound:\n%s", rust)
64+
}
65+
if !strings.Contains(rust, "Vec<Rc<RefCell<Option<N>>>>") {
66+
t.Fatalf("slice of interface-constrained type parameter should use wrapped elements:\n%s", rust)
67+
}
68+
if strings.Contains(rust, "Vec<N>") {
69+
t.Fatalf("slice of interface-constrained type parameter should not emit unwrapped Vec<N>:\n%s", rust)
70+
}
71+
}
72+
3373
func TestTranspileConstDeclUsesPackageVisibility(t *testing.T) {
3474
var out strings.Builder
3575
TranspileConstDecl(&out, &ast.GenDecl{

‎go/types.go‎

Lines changed: 31 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -217,6 +217,34 @@ func goTypeParamConstraintToRust(t types.Type) (string, bool) {
217217
return "", false
218218
}
219219

220+
func goTypeParamTraitConstraintName(t types.Type) (string, bool) {
221+
tp, ok := types.Unalias(t).(*types.TypeParam)
222+
if !ok || tp.Constraint() == nil {
223+
return "", false
224+
}
225+
named, ok := types.Unalias(tp.Constraint()).(*types.Named)
226+
if !ok || named.Obj() == nil {
227+
return "", false
228+
}
229+
iface, ok := named.Underlying().(*types.Interface)
230+
if !ok || iface.NumMethods() == 0 {
231+
return "", false
232+
}
233+
return goTypesNamedTypeToRust(named), true
234+
}
235+
236+
func goTypeParamTraitConstraintNameFromExpr(expr ast.Expr) (string, bool) {
237+
typeInfo := GetTypeInfo()
238+
if typeInfo == nil {
239+
return "", false
240+
}
241+
typ := typeInfo.GetType(expr)
242+
if typ == nil {
243+
return "", false
244+
}
245+
return goTypeParamTraitConstraintName(typ)
246+
}
247+
220248
// getStructSignature creates a unique signature for a struct type based on its fields
221249
func getStructSignature(structType *ast.StructType) string {
222250
var sig strings.Builder
@@ -862,6 +890,9 @@ func goCollectionElemTypeToRust(expr ast.Expr) string {
862890
if ident, ok := expr.(*ast.Ident); ok && ident.Name == "error" {
863891
return GoTypeToRust(expr)
864892
}
893+
if _, ok := goTypeParamTraitConstraintNameFromExpr(expr); ok {
894+
return GoTypeToRust(expr)
895+
}
865896
// Local named interfaces need wrapped slice elements: Box<dyn Trait> has no
866897
// Default (so make([]Trait, n) breaks) and can't be nil (so
867898
// interface-typed slice elements can't represent Go's nil interface value).

0 commit comments

Comments
 (0)