diff --git a/cmd/geth/main.go b/cmd/geth/main.go index 671d06ae2d..8e434948e2 100644 --- a/cmd/geth/main.go +++ b/cmd/geth/main.go @@ -96,6 +96,7 @@ func init() { utils.EthashDatasetsOnDiskFlag, utils.FastSyncFlag, utils.LightModeFlag, + utils.SyncModeFlag, utils.LightServFlag, utils.LightPeersFlag, utils.LightKDFFlag, diff --git a/cmd/utils/customflags.go b/cmd/utils/customflags.go index 00f28f2ecf..e5bf8724c1 100644 --- a/cmd/utils/customflags.go +++ b/cmd/utils/customflags.go @@ -17,6 +17,7 @@ package utils import ( + "encoding" "errors" "flag" "fmt" @@ -78,6 +79,58 @@ func (self DirectoryFlag) Apply(set *flag.FlagSet) { }) } +type TextMarshaler interface { + encoding.TextMarshaler + encoding.TextUnmarshaler +} + +// textMarshalerVal turns a TextMarshaler into a flag.Value +type textMarshalerVal struct { + v TextMarshaler +} + +func (v textMarshalerVal) String() string { + if v.v == nil { + return "" + } + text, _ := v.v.MarshalText() + return string(text) +} + +func (v textMarshalerVal) Set(s string) error { + return v.v.UnmarshalText([]byte(s)) +} + +// TextMarshalerFlag wraps a TextMarshaler value. +type TextMarshalerFlag struct { + Name string + Value TextMarshaler + Usage string +} + +func (f TextMarshalerFlag) GetName() string { + return f.Name +} + +func (f TextMarshalerFlag) String() string { + return fmt.Sprintf("%s \"%v\"\t%v", prefixedNames(f.Name), f.Value, f.Usage) +} + +func (f TextMarshalerFlag) Apply(set *flag.FlagSet) { + eachName(f.Name, func(name string) { + set.Var(textMarshalerVal{f.Value}, f.Name, f.Usage) + }) +} + +// GlobalTextMarshaler returns the value of a TextMarshalerFlag from the global flag set. +func GlobalTextMarshaler(ctx *cli.Context, name string) TextMarshaler { + val := ctx.GlobalGeneric(name) + if val == nil { + return nil + } + return val.(textMarshalerVal).v +} + // BigFlag is a command line flag that accepts 256 bit big integers in decimal or // hexadecimal syntax. type BigFlag struct { diff --git a/cmd/utils/flags.go b/cmd/utils/flags.go index 694eb21ea1..b4c7b88bde 100644 --- a/cmd/utils/flags.go +++ b/cmd/utils/flags.go @@ -174,6 +174,13 @@ var ( Name: "light", Usage: "Enable light client mode", } + defaultSyncMode = eth.DefaultConfig.SyncMode + SyncModeFlag = TextMarshalerFlag{ + Name: "syncmode", + Usage: `Blockchain sync mode ("fast", "full", or "light")`, + Value: &defaultSyncMode, + } + LightServFlag = cli.IntFlag{ Name: "lightserv", Usage: "Maximum percentage of time allowed for serving LES requests (0-90)", @@ -760,11 +767,11 @@ func setEthash(ctx *cli.Context, cfg *eth.Config) { } } -func checkExclusive(ctx *cli.Context, flags ...cli.BoolFlag) { +func checkExclusive(ctx *cli.Context, flags ...cli.Flag) { set := make([]string, 0, 1) for _, flag := range flags { - if ctx.GlobalIsSet(flag.Name) { - set = append(set, "--"+flag.Name) + if ctx.GlobalIsSet(flag.GetName()) { + set = append(set, "--"+flag.GetName()) } } if len(set) > 1 { @@ -776,16 +783,19 @@ func checkExclusive(ctx *cli.Context, flags ...cli.BoolFlag) { func SetEthConfig(ctx *cli.Context, stack *node.Node, cfg *eth.Config) { // Avoid conflicting network flags checkExclusive(ctx, DevModeFlag, TestNetFlag) - checkExclusive(ctx, FastSyncFlag, LightModeFlag) + checkExclusive(ctx, FastSyncFlag, LightModeFlag, SyncModeFlag) ks := stack.AccountManager().Backends(keystore.KeyStoreType)[0].(*keystore.KeyStore) setEtherbase(ctx, ks, cfg) setGPO(ctx, &cfg.GPO) setEthash(ctx, cfg) - if ctx.GlobalBool(FastSyncFlag.Name) { + switch { + case ctx.GlobalIsSet(SyncModeFlag.Name): + cfg.SyncMode = *GlobalTextMarshaler(ctx, SyncModeFlag.Name).(*downloader.SyncMode) + case ctx.GlobalBool(FastSyncFlag.Name): cfg.SyncMode = downloader.FastSync - } else if ctx.GlobalBool(LightModeFlag.Name) { + case ctx.GlobalBool(LightModeFlag.Name): cfg.SyncMode = downloader.LightSync } if ctx.GlobalIsSet(LightServFlag.Name) {