evcc-io/cmd/decorate/decorate.go

347 lines
7.9 KiB
Go

package main
import (
"bufio"
"bytes"
_ "embed"
"fmt"
"go/format"
"io"
"os"
"reflect"
"strconv"
"strings"
"text/template"
"github.com/Masterminds/sprig/v3"
"github.com/evcc-io/evcc/api"
"github.com/spf13/pflag"
"golang.org/x/tools/imports"
)
//go:generate go tool decorate
//evcc:function decorateTest
//evcc:basetype api.Charger
//evcc:types api.MeterEnergy,api.PhaseSwitcher,api.PhaseGetter
//go:embed decorate.tpl
var srcTmpl string
//go:embed header.tpl
var header string
type funcStruct struct {
Signature, Function, VarName, ReturnTypes, BaseType, ShortType string
Params []string
}
type typeStruct struct {
Type, ShortType string
Functions []funcStruct
}
var interfaces = make(map[string]reflect.Type)
func init() {
for _, typ := range []reflect.Type{
reflect.TypeFor[api.BatteryCapacity](),
reflect.TypeFor[api.SocLimiter](),
reflect.TypeFor[api.BatteryController](),
reflect.TypeFor[api.BatterySocLimiter](),
reflect.TypeFor[api.BatteryPowerLimiter](),
reflect.TypeFor[api.PhasePowers](),
reflect.TypeFor[api.PhaseGetter](),
reflect.TypeFor[api.CurrentController](),
reflect.TypeFor[api.ChargeController](),
reflect.TypeFor[api.CurrentController](),
reflect.TypeFor[api.PhaseCurrents](),
reflect.TypeFor[api.PhaseSwitcher](),
reflect.TypeFor[api.Battery](),
reflect.TypeFor[api.ChargeState](),
reflect.TypeFor[api.MeterEnergy](),
reflect.TypeFor[api.PhaseCurrents](),
reflect.TypeFor[api.PhaseVoltages](),
reflect.TypeFor[api.MaxACPowerGetter](),
reflect.TypeFor[api.Meter](),
reflect.TypeFor[api.CurrentGetter](),
reflect.TypeFor[api.Curtailer](),
reflect.TypeFor[api.Resurrector](),
reflect.TypeFor[api.VehicleOdometer](),
reflect.TypeFor[api.VehicleRange](),
reflect.TypeFor[api.VehicleClimater](),
reflect.TypeFor[api.VehicleFinishTimer](),
reflect.TypeFor[api.Identifier](),
reflect.TypeFor[api.ChargerEx](),
reflect.TypeFor[api.ChargeRater](),
reflect.TypeFor[api.StatusReasoner](),
} {
interfaces[typ.String()] = typ
}
}
func getTemplate(dtypes []reflect.Type, types map[string]typeStruct) *template.Template {
tmpl, err := template.New("gen").Funcs(sprig.FuncMap()).Funcs(template.FuncMap{
// orderedParams returns a slice of funcStruct ordered by dynamicType
"orderedParams": func() []funcStruct {
orderedParams := make([]funcStruct, 0)
for _, t := range dtypes {
for _, f := range types[getTypeImport(t)].Functions {
orderedParams = append(orderedParams, f)
}
}
return orderedParams
},
}).Parse(srcTmpl)
if err != nil {
fmt.Printf("invalid template: %s", err)
os.Exit(2)
}
return tmpl
}
func getTypeImport(t reflect.Type) string {
n := t.Name()
if p := t.PkgPath(); p != "" {
if s := strings.Split(p, "github.com/evcc-io/evcc/"); len(s) == 2 {
return fmt.Sprintf("%s.%s", s[1], n)
} else {
return fmt.Sprintf("%s.%s", p, n)
}
}
return n
}
func generate(out io.Writer, functionName, baseType string, dtypes []reflect.Type) error {
types := make(map[string]typeStruct)
for _, t := range dtypes {
var funcs []funcStruct
lastPart := t.Name()
for i := 0; i < t.NumMethod(); i++ {
m := t.Method(i)
varName := strings.ToLower(lastPart[:1]) + lastPart[1:]
if t.NumMethod() > 1 {
varName += strconv.Itoa(i)
}
var params []string
for input := range m.Type.Ins() {
params = append(params, getTypeImport(input))
}
var returns []string
for output := range m.Type.Outs() {
returns = append(returns, getTypeImport(output))
}
funcs = append(funcs, funcStruct{
VarName: varName,
Signature: fmt.Sprintf("func(%s) (%s)", strings.Join(params, ", "), strings.Join(returns, ", ")),
Function: m.Name,
Params: params,
ReturnTypes: fmt.Sprintf("(%s)", strings.Join(returns, ",")),
BaseType: t.String(),
ShortType: t.Name(),
})
}
types[getTypeImport(t)] = typeStruct{
Type: t.Name(),
ShortType: lastPart,
Functions: funcs,
}
}
returnType := *ret
if returnType == "" {
returnType = baseType
}
shortBase := strings.TrimLeft(baseType, "*")
if baseTypeParts := strings.SplitN(baseType, ".", 2); len(baseTypeParts) > 1 {
shortBase = baseTypeParts[1]
}
vars := struct {
Function string
BaseType, ShortBase string
ReturnType string
Types map[string]typeStruct
Combinations [][]string
}{
Function: functionName,
BaseType: baseType,
ShortBase: shortBase,
ReturnType: returnType,
Types: types,
}
return getTemplate(dtypes, types).Execute(out, vars)
}
type decorationSet struct {
function, base, ret, types string
}
var (
target = pflag.StringP("out", "o", "", "output file")
pkg = pflag.StringP("package", "p", "", "package name")
funcname = pflag.StringP("function", "f", "", "function name")
base = pflag.StringP("base", "b", "", "base type")
ret = pflag.StringP("return", "r", "", "return type")
types = pflag.StringP("type", "t", "", "comma-separated list of type definitions")
)
// Usage prints flags usage
func Usage() {
fmt.Fprintf(os.Stderr, "Usage of decorate:\n")
fmt.Fprintf(os.Stderr, "\ndecorate [flags] -type interface,interface function,function signature\n")
fmt.Fprintf(os.Stderr, "\nFlags:\n")
pflag.PrintDefaults()
}
func parseFile(file string) ([]decorationSet, error) {
var res []decorationSet
f, err := os.Open(file)
if err != nil {
return nil, err
}
defer f.Close()
var current decorationSet
scanner := bufio.NewScanner(f)
for scanner.Scan() {
line := scanner.Text()
if s, ok := strings.CutPrefix(line, "//evcc:"); ok {
segs := strings.SplitN(s, " ", 2)
if len(segs) != 2 {
panic("invalid segments: " + s)
}
switch segs[0] {
case "function":
// must be first
if current.function != "" {
res = append(res, current)
current = decorationSet{}
}
current.function = segs[1]
case "basetype":
current.base = segs[1]
case "returntype":
current.ret = segs[1]
case "types":
current.types = segs[1]
default:
panic("invalid directive //evcc:" + segs[0])
}
}
}
if current.function != "" {
res = append(res, current)
}
return res, scanner.Err()
}
func main() {
pflag.Usage = Usage
pflag.Parse()
// read target from go:generate
gofile, ok := os.LookupEnv("GOFILE")
if *target == "" && ok {
gofile := strings.TrimSuffix(gofile, ".go") + "_decorators.go"
target = &gofile
}
// read target from go:generate
if gopkg, ok := os.LookupEnv("GOPACKAGE"); *pkg == "" && ok {
pkg = &gopkg
}
sets := []decorationSet{{*funcname, *base, *ret, *types}}
if *funcname == "" {
all, err := parseFile(gofile)
if err != nil {
fmt.Println(err)
os.Exit(2)
}
sets = all
}
if *pkg == "" || len(sets) == 0 || sets[0].base == "" || len(sets[0].types) == 0 {
Usage()
os.Exit(2)
}
var out io.Writer = os.Stdout
var name string
if target != nil {
name = *target
if !strings.HasSuffix(name, ".go") {
name += ".go"
}
dst, err := os.Create(name)
if err != nil {
fmt.Println(err)
os.Exit(2)
}
defer dst.Close()
out = dst
}
generated := new(bytes.Buffer)
fmt.Fprintln(generated, strings.ReplaceAll(header, "{{.Package}}", *pkg))
for _, set := range sets {
var types []reflect.Type
for t := range strings.SplitSeq(set.types, ",") {
typ, ok := interfaces[t]
if !ok {
fmt.Printf("don't know interface %s\n", t)
os.Exit(2)
}
types = append(types, typ)
}
var buf bytes.Buffer
if err := generate(&buf, set.function, set.base, types); err != nil {
fmt.Println(err)
os.Exit(2)
}
fmt.Fprintln(generated, buf.String())
}
formatted, err := format.Source(generated.Bytes())
if err != nil {
formatted = generated.Bytes()
}
formatted, err = imports.Process(name, formatted, nil)
if err != nil {
fmt.Println(err)
os.Exit(3)
}
if _, err := out.Write(formatted); err != nil {
fmt.Println(err)
os.Exit(2)
}
}