make protoc generation compatible with api v2+ (#9673)

Starting with the Teleport 9 release, we will be versioning the
API module. This change ensures that the generated protobuf code
imports the correct version of the API by:

- introducing a small new command to print the correct version
- adding import rewrite rules to the protoc invocation
This commit is contained in:
Brian Joerger
2022-01-24 19:16:05 +00:00
committed by GitHub
parent fdf921fb50
commit eb40cdc73e
5 changed files with 204 additions and 89 deletions
+32 -26
View File
@@ -642,7 +642,7 @@ version: $(VERSRC)
$(VERSRC): Makefile
VERSION=$(VERSION) $(MAKE) -f version.mk setver
# Update api module path, but don't fail on error.
$(MAKE) update-api-module-path || true
$(MAKE) update-api-import-path || true
# This rule updates the api module path to be in sync with the current api release version.
# e.g. github.com/gravitational/teleport/api/vX -> github.com/gravitational/teleport/api/vY
@@ -654,11 +654,10 @@ $(VERSRC): Makefile
# - v0.0.0 -> v1.0.0 - both have no version suffix - github.com/gravitational/teleport/api
#
# Note: any build flags needed to compile go files (such as build tags) should be provided below.
.PHONY: update-api-module-path
update-api-module-path:
# update-api-module-path is temporarily disabled because currently `make grpc` does not know how to deal with v2+ go modules.
# go run build.assets/update_api_module_path/main.go -tags "bpf fips pam roletester desktop_access_rdp"
# $(MAKE) grpc
.PHONY: update-api-import-path
update-api-import-path:
go run build.assets/gomod/update-api-import-path/main.go -tags "bpf fips pam roletester desktop_access_rdp linux"
$(MAKE) grpc
# make tag - prints a tag to use with git for the current version
# To put a new release on Github:
@@ -734,10 +733,18 @@ enter:
grpc:
$(MAKE) -C build.assets grpc
# proto file dependencies within the api module must be passed with the 'M' flag. This
# way protoc generated files will use the correct api module import path in the case where
# the import path has a version suffix, e.g. github.com/gravitational/teleport/api/v8
GOGOPROTO_IMPORTMAP ?= $\
Mgithub.com/gravitational/teleport/api/types/types.proto=$(API_IMPORT_PATH)/types,$\
Mgithub.com/gravitational/teleport/api/types/events/events.proto=$(API_IMPORT_PATH)/types/events,$\
Mgithub.com/gravitational/teleport/api/types/wrappers/wrappers.proto=$(API_IMPORT_PATH)/types/wrappers,$\
Mgithub.com/gravitational/teleport/api/types/webauthn/webauthn.proto=$(API_IMPORT_PATH)/types/webauthn
# buildbox-grpc generates GRPC stubs
.PHONY: buildbox-grpc
buildbox-grpc:
# standard GRPC output
echo $$PROTO_INCLUDE
$(CLANG_FORMAT) -i -style='{ColumnLimit: 100, IndentWidth: 4, Language: Proto}' \
api/client/proto/authservice.proto \
@@ -750,47 +757,46 @@ buildbox-grpc:
lib/multiplexer/test/ping.proto \
lib/web/envelope.proto
protoc -I=.:$$PROTO_INCLUDE \
--proto_path=api/client/proto \
--gogofast_out=plugins=grpc:api/client/proto \
# we eval within the make target to avoid invoking `go run` with every other call to the makefile
$(eval API_IMPORT_PATH := $(shell go run build.assets/gomod/print-import-path/main.go ./api))
cd api/client/proto && protoc -I=.:$$PROTO_INCLUDE \
--gogofast_out=plugins=grpc,$(GOGOPROTO_IMPORTMAP):. \
authservice.proto
protoc -I=.:$$PROTO_INCLUDE \
--proto_path=api/types/events \
--gogofast_out=plugins=grpc:api/types/events \
cd api/types/events && protoc -I=.:$$PROTO_INCLUDE \
--gogofast_out=plugins=grpc,$(GOGOPROTO_IMPORTMAP):. \
events.proto
protoc -I=.:$$PROTO_INCLUDE \
--proto_path=api/types \
--gogofast_out=plugins=grpc:api/types \
cd api/types && protoc -I=.:$$PROTO_INCLUDE \
--gogofast_out=plugins=grpc,$(GOGOPROTO_IMPORTMAP):. \
types.proto
protoc -I=.:$$PROTO_INCLUDE \
--proto_path=api/types/webauthn \
--gogofast_out=plugins=grpc:api/types/webauthn \
cd api/types/webauthn && protoc -I=.:$$PROTO_INCLUDE \
--gogofast_out=plugins=grpc,$(GOGOPROTO_IMPORTMAP):. \
webauthn.proto
protoc -I=.:$$PROTO_INCLUDE \
--proto_path=api/types/wrappers \
--gogofast_out=plugins=grpc:api/types/wrappers \
cd api/types/wrappers && protoc -I=.:$$PROTO_INCLUDE \
--gogofast_out=plugins=grpc,$(GOGOPROTO_IMPORTMAP):. \
wrappers.proto
cd lib/datalog && protoc -I=.:$$PROTO_INCLUDE \
--gogofast_out=plugins=grpc:. \
--gogofast_out=plugins=grpc,$(GOGOPROTO_IMPORTMAP):. \
types.proto
cd lib/events && protoc -I=.:$$PROTO_INCLUDE \
--gogofast_out=plugins=grpc:. \
--gogofast_out=plugins=grpc,$(GOGOPROTO_IMPORTMAP):. \
slice.proto
cd lib/multiplexer/test && protoc -I=.:$$PROTO_INCLUDE \
--gogofast_out=plugins=grpc:. \
--gogofast_out=plugins=grpc,$(GOGOPROTO_IMPORTMAP):. \
ping.proto
cd lib/web && protoc -I=.:$$PROTO_INCLUDE \
--gogofast_out=plugins=grpc:. \
--gogofast_out=plugins=grpc,$(GOGOPROTO_IMPORTMAP):. \
envelope.proto
.PHONY: goinstall
goinstall:
go install $(BUILDFLAGS) \
+40
View File
@@ -0,0 +1,40 @@
// Copyright 2022 Gravitational, Inc
//
// Licensed under the Apache License, Version 2.0 (the "License");
// you may not use this file except in compliance with the License.
// You may obtain a copy of the License at
//
// http://www.apache.org/licenses/LICENSE-2.0
//
// Unless required by applicable law or agreed to in writing, software
// distributed under the License is distributed on an "AS IS" BASIS,
// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
// See the License for the specific language governing permissions and
// limitations under the License.
package gomod
import (
"os"
"path/filepath"
"github.com/gravitational/trace"
"golang.org/x/mod/modfile"
)
// GetImportPath gets the module's import path from its go.mod file
func GetImportPath(dir string) (string, error) {
modPath := filepath.Join(dir, "go.mod")
bts, err := os.ReadFile(modPath)
if err != nil {
return "", trace.Wrap(err)
}
modFile, err := modfile.Parse(modPath, bts, nil /* fix */)
if err != nil {
return "", trace.Wrap(err)
}
if modFile.Module == nil || modFile.Module.Mod.Path == "" {
return "", trace.NotFound("could not find mod path for %v", dir)
}
return modFile.Module.Mod.Path, nil
}
@@ -0,0 +1,40 @@
/*
Copyright 2022 Gravitational, Inc.
Licensed under the Apache License, Version 2.0 (the "License");
you may not use this file except in compliance with the License.
You may obtain a copy of the License at
http://www.apache.org/licenses/LICENSE-2.0
Unless required by applicable law or agreed to in writing, software
distributed under the License is distributed on an "AS IS" BASIS,
WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
See the License for the specific language governing permissions and
limitations under the License.
*/
// Command print-import-path prints the import path that
// should appear in Go import paths to stdout.
package main
import (
"fmt"
"log"
"os"
"github.com/gravitational/teleport/build.assets/gomod"
"github.com/gravitational/trace"
)
// prints the import path of the api module
func main() {
if len(os.Args) < 1 {
log.Fatal("first argument should be a path to a go.mod file")
}
goModFilePath := os.Args[1]
modPath, err := gomod.GetImportPath(goModFilePath)
if err != nil {
log.Fatal(trace.Wrap(err))
}
fmt.Println(modPath)
}
@@ -11,6 +11,11 @@ See the License for the specific language governing permissions and
limitations under the License.
*/
// Command update-api-import-path updates the api import path to
// incorporate the version set in /api/version.go. If the major
// version hasn't changed or the version is a prelease, no change
// is made. Otherwise, all go.mod files, .go files, and .proto files
// are updated to use the new api import path as needed.
package main
import (
@@ -22,10 +27,11 @@ import (
"path/filepath"
"strings"
"github.com/coreos/go-semver/semver"
"github.com/gravitational/teleport/api"
"github.com/gravitational/teleport/build.assets/gomod"
"github.com/gravitational/teleport/lib/utils"
"github.com/coreos/go-semver/semver"
"github.com/gravitational/trace"
log "github.com/sirupsen/logrus"
"golang.org/x/mod/modfile"
@@ -37,6 +43,8 @@ func init() {
utils.InitLogger(utils.LoggingForCLI, log.DebugLevel)
}
// This script should only be run through the make target `make update-api-module-path`
// since it relies on relative paths to the /api and root directories.
func main() {
var buildFlags []string
if len(os.Args) > 1 {
@@ -45,13 +53,13 @@ func main() {
}
// the api module import path should only be updated on releases
newVersion := api.Version
if isPreRelease(newVersion) {
newVersion := semver.New(api.Version)
if newVersion.PreRelease != "" {
exitWithMessage("the current API version (%v) is not a release, continue without updating", newVersion)
}
// get the current api module import path
currentModPath, err := getModImportPath("./api")
currentModPath, err := gomod.GetImportPath("./api")
if err != nil {
exitWithError(trace.Wrap(err, "failed to get current mod path"), nil)
}
@@ -67,11 +75,11 @@ func main() {
// update go files within the teleport/api and teleport modules to use the new import path
log.Info("Updating teleport/api module...")
if err := updateGoModule("./api", currentModPath, newPath, newVersion, buildFlags, addRollBack); err != nil {
if err := updateGoModule("./api", currentModPath, newPath, newVersion.String(), buildFlags, addRollBack); err != nil {
exitWithError(trace.Wrap(err, "failed to update teleport/api module"), rollBackFuncs)
}
log.Info("Updating teleport module...")
if err := updateGoModule("./", currentModPath, newPath, newVersion, buildFlags, addRollBack); err != nil {
if err := updateGoModule("./", currentModPath, newPath, newVersion.String(), buildFlags, addRollBack); err != nil {
exitWithError(trace.Wrap(err, "failed to update teleport module"), rollBackFuncs)
}
@@ -123,12 +131,14 @@ func updateGoImports(p *packages.Package, currentPath, newPath string, addRollBa
var rewritten bool
for _, i := range syn.Imports {
imp := strings.Replace(i.Path.Value, "\"", "", 2)
if strings.HasPrefix(imp, currentPath) && !strings.HasPrefix(imp, newPath) {
// Replace all instances of the current path with the new path, but prevent this from happening multiple times in edge cases
if strings.HasPrefix(imp, currentPath) && (!strings.HasPrefix(newPath, currentPath) || !strings.HasPrefix(imp, newPath)) {
newImp := strings.Replace(imp, currentPath, newPath, 1)
if astutil.RewriteImport(p.Fset, syn, imp, newImp) {
rewritten = true
}
}
}
if !rewritten {
continue
@@ -231,8 +241,8 @@ func updateGoModFile(dir, oldPath, newPath, newVersion string, addRollBackFunc a
return nil
}
// updateProtoFiles updates instances of the currentPath with
// the newPath in .proto files within the given directory
// updateProtoFiles updates gogoproto cast types and custom types in .proto files
// within the given directory to use the new api import path.
func updateProtoFiles(rootDir, currentPath, newPath string, addRollBackFunc addRollBackFunc) error {
return filepath.WalkDir(rootDir, func(path string, d fs.DirEntry, err error) error {
if err != nil {
@@ -244,7 +254,16 @@ func updateProtoFiles(rootDir, currentPath, newPath string, addRollBackFunc addR
return trace.Wrap(err)
}
updatedData := bytes.ReplaceAll(data, []byte(currentPath), []byte(newPath))
// Replace all instances of the api import path in gogoproto casttypes with the new import path
currentCastTypes := fmt.Sprintf(`(gogoproto.casttype) = "%v`, currentPath)
newCastTypes := fmt.Sprintf(`(gogoproto.casttype) = "%v`, newPath)
updatedData := bytes.ReplaceAll(data, []byte(currentCastTypes), []byte(newCastTypes))
// Replace all instances of the api import path in gogoproto customtypes with the new import path
currentCustomTypes := fmt.Sprintf(`(gogoproto.customtype) = "%v`, currentPath)
newCustomTypes := fmt.Sprintf(`(gogoproto.customtype) = "%v`, newPath)
updatedData = bytes.ReplaceAll(updatedData, []byte(currentCustomTypes), []byte(newCustomTypes))
fileMode := d.Type().Perm()
if err := os.WriteFile(path, updatedData, fileMode); err != nil {
return trace.Wrap(err)
@@ -259,38 +278,12 @@ func updateProtoFiles(rootDir, currentPath, newPath string, addRollBackFunc addR
})
}
// getModImportPath gets the module's currently set path/name
func getModImportPath(dir string) (string, error) {
modFile, err := getModFile(dir)
if err != nil {
return "", trace.Wrap(err)
}
if modFile.Module.Mod.Path == "" {
return "", trace.NotFound("could not find mod path for %v", dir)
}
return modFile.Module.Mod.Path, nil
}
// getModFile returns an AST of the given go.mod file
func getModFile(dir string) (*modfile.File, error) {
modPath := filepath.Join(dir, "go.mod")
bts, err := os.ReadFile(modPath)
if err != nil {
return nil, trace.Wrap(err)
}
f, err := modfile.Parse(modPath, bts, nil)
if err != nil {
return nil, trace.Wrap(err)
}
return f, nil
}
// getNewModImportPath gets the new import path given a go module import path and the updated version
func getNewModImportPath(oldPath, newVersion string) string {
func getNewModImportPath(oldPath string, newVersion *semver.Version) string {
// get the new major version suffix - e.g "" for v0/v1 or "/vX" for vX where X >= 2
var majVerSuffix string
if ver := semver.New(newVersion); ver.Major >= 2 {
majVerSuffix = fmt.Sprintf("/v%d", ver.Major)
if newVersion.Major >= 2 {
majVerSuffix = fmt.Sprintf("/v%d", newVersion.Major)
}
// get the new mod path by replacing the current mod path with the new major version suffix
@@ -301,11 +294,6 @@ func getNewModImportPath(oldPath, newVersion string) string {
return newPath
}
// returns whether the current api version is a pre-release, e.g "v7.0.0-beta"
func isPreRelease(version string) bool {
return semver.New(version).PreRelease != ""
}
// rollBackFuncs are used to revert changes if the program fails with an error.
type rollBackFunc func() error
type addRollBackFunc func(rollBackFunc)
@@ -21,6 +21,9 @@ import (
"strings"
"testing"
"github.com/gravitational/teleport/build.assets/gomod"
"github.com/coreos/go-semver/semver"
"github.com/stretchr/testify/require"
)
@@ -36,43 +39,81 @@ import (
"other/mod/path"
)
`
updatedGoFile := `package main
import "mod/path/v2"
import "other/mod/path"
import (
"mod/path/v2"
alias "mod/path/v2"
"other/mod/path"
)
`
// Create a dummy go module with go.mod and main.go file
pkgDir := t.TempDir()
writeFile(t, pkgDir, "go.mod", newGoModFileString("pkg"))
goFilePath := writeFile(t, pkgDir, "main.go", goFile)
addRollBack := testRollBack(t, goFilePath, goFile)
// Run main.go file through the update function
err := updateGoPkgs(pkgDir, "mod/path", "updated/mod/path", nil, addRollBack)
goFilePath := writeFile(t, pkgDir, "main.go", goFile)
addRollBack := testRollBack(t, goFilePath, goFile)
err := updateGoPkgs(pkgDir, "mod/path", "mod/path/v2", nil, addRollBack)
require.NoError(t, err)
readAndCompareFile(t, goFilePath, updatedGoFile)
// Read the updated file and expect all instances of "mod/path" to be replaced with "updated/mod/path"
readAndCompareFile(t, goFilePath, strings.ReplaceAll(goFile, "\"mod/path\"", "\"updated/mod/path\""))
// Run updated main.go file through update function
goFilePath = writeFile(t, pkgDir, "main.go", updatedGoFile)
addRollBack = testRollBack(t, goFilePath, updatedGoFile)
err = updateGoPkgs(pkgDir, "mod/path/v2", "mod/path", nil, addRollBack)
require.NoError(t, err)
readAndCompareFile(t, goFilePath, goFile)
}
func TestUpdateProtoFiles(t *testing.T) {
protoFile := `syntax = "proto3";
package proto;
import "mod/path/types.proto";
message Example {
types.Type field = 6 [
types.Type field1 = 1 [
(gogoproto.casttype) = "mod/path/types.Traits"
];
types.Type field2 = 2 [
(gogoproto.customtype) = "mod/path/types.Traits"
];
}
`
// Only update casttype and customtype options.
updatedProtoFile := `syntax = "proto3";
package proto;
import "mod/path/types.proto";
message Example {
types.Type field1 = 1 [
(gogoproto.casttype) = "mod/path/v2/types.Traits"
];
types.Type field2 = 2 [
(gogoproto.customtype) = "mod/path/v2/types.Traits"
];
}
`
// Write proto file to disk
dir := t.TempDir()
protoFilePath := writeFile(t, dir, "proto.proto", protoFile)
addRollBack := testRollBack(t, protoFilePath, protoFile)
// Run proto file through update function
err := updateProtoFiles(dir, "mod/path", "updated/mod/path", addRollBack)
protoFilePath := writeFile(t, dir, "proto.proto", protoFile)
addRollBack := testRollBack(t, protoFilePath, protoFile)
err := updateProtoFiles(dir, "mod/path", "mod/path/v2", addRollBack)
require.NoError(t, err)
readAndCompareFile(t, protoFilePath, updatedProtoFile)
// Read the updated file and expect all instances of "mod/path" to be replaced with "updated/mod/path"
readAndCompareFile(t, protoFilePath, strings.ReplaceAll(protoFile, "mod/path", "updated/mod/path"))
// Run updated proto file through update function
protoFilePath = writeFile(t, dir, "proto.proto", updatedProtoFile)
addRollBack = testRollBack(t, protoFilePath, updatedProtoFile)
err = updateProtoFiles(dir, "mod/path/v2", "mod/path", addRollBack)
require.NoError(t, err)
readAndCompareFile(t, protoFilePath, protoFile)
}
func TestUpdateGoModulePath(t *testing.T) {
@@ -98,12 +139,12 @@ func TestUpdateGoModulePath(t *testing.T) {
newGoModFileString("updated/go/mod/header"),
))
t.Run("updated module in statements", testUpdate("mod/path", "updated/mod/path", "1.2.3",
t.Run("updated module in statements", testUpdate("mod/path", "mod/path/v2", "1.2.3",
newGoModFileString("go/mod/header", requireStatement("mod/path", "0.1.2")),
newGoModFileString("go/mod/header", requireStatement("updated/mod/path", "1.2.3")),
newGoModFileString("go/mod/header", requireStatement("mod/path/v2", "1.2.3")),
))
t.Run("updated module not in mod file", testUpdate("mod/path", "updated/mod/path", "1.2.3",
t.Run("updated module not in mod file", testUpdate("mod/path", "mod/path/v2", "1.2.3",
newGoModFileString("go/mod/header", requireStatement("other/mod/path", "0.1.2")),
newGoModFileString("go/mod/header", requireStatement("other/mod/path", "0.1.2")),
))
@@ -112,7 +153,7 @@ func TestUpdateGoModulePath(t *testing.T) {
// test that every statement and the header gets updated properly.
testUpdateAllStatements := func(oldModPath, oldVersion, newVersion string) func(*testing.T) {
oldModFile := newGoModFileString(oldModPath, allGoModStatements(oldModPath, oldVersion)...)
newModPath := getNewModImportPath(oldModPath, newVersion)
newModPath := getNewModImportPath(oldModPath, semver.New(newVersion))
newModFile := newGoModFileString(newModPath, allGoModStatements(newModPath, newVersion)...)
return testUpdate(oldModPath, newModPath, newVersion, oldModFile, newModFile)
}
@@ -134,9 +175,9 @@ func TestGetImportPaths(t *testing.T) {
writeFile(t, modDir, "go.mod", newGoModFileString(currentModPath))
// Get import paths using the mod file in disk
oldModPath, err := getModImportPath(modDir)
oldModPath, err := gomod.GetImportPath(modDir)
require.NoError(t, err)
newModPath := getNewModImportPath(oldModPath, newVersion)
newModPath := getNewModImportPath(oldModPath, semver.New(newVersion))
// Compare paths to expected results
require.Equal(t, currentModPath, oldModPath)