Compare commits
	
		
			32 Commits
		
	
	
		
	
	| Author | SHA1 | Date | |
|---|---|---|---|
| ae952b2166 | |||
| b24dba9a45 | |||
| cfbc20367d | |||
| e25912758e | |||
| e1ae77a9db | |||
| 9d07b3955f | |||
| 02be696c25 | |||
| ba07625b7c | |||
| aeded3fb37 | |||
| 1a1cd6d0aa | |||
| 64cc1342a0 | |||
| 8431b6adf5 | |||
| 24e923fe84 | |||
| 10ddc7c190 | |||
| 7f88a0726c | |||
| 2224db8e85 | |||
| c60afc89bb | |||
| bbb33e9fd6 | |||
| ac05eff1e8 | |||
| 1aaad66233 | |||
| d4994b8c8d | |||
| e3b8d2cc0f | |||
| fff609db4a | |||
| 5e99e07f40 | |||
| bdb181cb3a | |||
| 3552acd38b | |||
| c42324c58f | |||
| 3a9c3f4e9e | |||
| becd8f1ebc | |||
| e733f30c38 | |||
| 1a9e5c70fc | |||
| f3700a772d | 
							
								
								
									
										142
									
								
								cmdext/cmdrunner.go
									
									
									
									
									
										Normal file
									
								
							
							
						
						
									
										142
									
								
								cmdext/cmdrunner.go
									
									
									
									
									
										Normal file
									
								
							| @@ -0,0 +1,142 @@ | |||||||
|  | package cmdext | ||||||
|  |  | ||||||
|  | import ( | ||||||
|  | 	"bufio" | ||||||
|  | 	"os/exec" | ||||||
|  | 	"time" | ||||||
|  | ) | ||||||
|  |  | ||||||
|  | type CommandResult struct { | ||||||
|  | 	StdOut          string | ||||||
|  | 	StdErr          string | ||||||
|  | 	StdCombined     string | ||||||
|  | 	ExitCode        int | ||||||
|  | 	CommandTimedOut bool | ||||||
|  | } | ||||||
|  |  | ||||||
|  | func RunCommand(program string, args []string, timeout *time.Duration) (CommandResult, error) { | ||||||
|  |  | ||||||
|  | 	cmd := exec.Command(program, args...) | ||||||
|  |  | ||||||
|  | 	stdoutPipe, err := cmd.StdoutPipe() | ||||||
|  | 	if err != nil { | ||||||
|  | 		return CommandResult{}, err | ||||||
|  | 	} | ||||||
|  |  | ||||||
|  | 	stderrPipe, err := cmd.StderrPipe() | ||||||
|  | 	if err != nil { | ||||||
|  | 		return CommandResult{}, err | ||||||
|  | 	} | ||||||
|  |  | ||||||
|  | 	err = cmd.Start() | ||||||
|  | 	if err != nil { | ||||||
|  | 		return CommandResult{}, err | ||||||
|  | 	} | ||||||
|  |  | ||||||
|  | 	errch := make(chan error, 1) | ||||||
|  | 	go func() { errch <- cmd.Wait() }() | ||||||
|  |  | ||||||
|  | 	combch := make(chan string, 32) | ||||||
|  | 	stopCombch := make(chan bool) | ||||||
|  |  | ||||||
|  | 	stdout := "" | ||||||
|  | 	go func() { | ||||||
|  | 		scanner := bufio.NewScanner(stdoutPipe) | ||||||
|  | 		for scanner.Scan() { | ||||||
|  | 			txt := scanner.Text() | ||||||
|  | 			stdout += txt | ||||||
|  | 			combch <- txt | ||||||
|  | 		} | ||||||
|  | 	}() | ||||||
|  |  | ||||||
|  | 	stderr := "" | ||||||
|  | 	go func() { | ||||||
|  | 		scanner := bufio.NewScanner(stderrPipe) | ||||||
|  | 		for scanner.Scan() { | ||||||
|  | 			txt := scanner.Text() | ||||||
|  | 			stderr += txt | ||||||
|  | 			combch <- txt | ||||||
|  | 		} | ||||||
|  | 	}() | ||||||
|  |  | ||||||
|  | 	defer func() { | ||||||
|  | 		stopCombch <- true | ||||||
|  | 	}() | ||||||
|  |  | ||||||
|  | 	stdcombined := "" | ||||||
|  | 	go func() { | ||||||
|  | 		for { | ||||||
|  | 			select { | ||||||
|  | 			case txt := <-combch: | ||||||
|  | 				stdcombined += txt | ||||||
|  | 			case <-stopCombch: | ||||||
|  | 				return | ||||||
|  | 			} | ||||||
|  | 		} | ||||||
|  | 	}() | ||||||
|  |  | ||||||
|  | 	if timeout != nil { | ||||||
|  |  | ||||||
|  | 		select { | ||||||
|  |  | ||||||
|  | 		case <-time.After(*timeout): | ||||||
|  | 			_ = cmd.Process.Kill() | ||||||
|  | 			return CommandResult{ | ||||||
|  | 				StdOut:          stdout, | ||||||
|  | 				StdErr:          stderr, | ||||||
|  | 				StdCombined:     stdcombined, | ||||||
|  | 				ExitCode:        -1, | ||||||
|  | 				CommandTimedOut: true, | ||||||
|  | 			}, nil | ||||||
|  |  | ||||||
|  | 		case err := <-errch: | ||||||
|  | 			if exiterr, ok := err.(*exec.ExitError); ok { | ||||||
|  | 				return CommandResult{ | ||||||
|  | 					StdOut:          stdout, | ||||||
|  | 					StdErr:          stderr, | ||||||
|  | 					StdCombined:     stdcombined, | ||||||
|  | 					ExitCode:        exiterr.ExitCode(), | ||||||
|  | 					CommandTimedOut: false, | ||||||
|  | 				}, nil | ||||||
|  | 			} else if err != nil { | ||||||
|  | 				return CommandResult{}, err | ||||||
|  | 			} else { | ||||||
|  | 				return CommandResult{ | ||||||
|  | 					StdOut:          stdout, | ||||||
|  | 					StdErr:          stderr, | ||||||
|  | 					StdCombined:     stdcombined, | ||||||
|  | 					ExitCode:        0, | ||||||
|  | 					CommandTimedOut: false, | ||||||
|  | 				}, nil | ||||||
|  | 			} | ||||||
|  |  | ||||||
|  | 		} | ||||||
|  |  | ||||||
|  | 	} else { | ||||||
|  |  | ||||||
|  | 		select { | ||||||
|  |  | ||||||
|  | 		case err := <-errch: | ||||||
|  | 			if exiterr, ok := err.(*exec.ExitError); ok { | ||||||
|  | 				return CommandResult{ | ||||||
|  | 					StdOut:          stdout, | ||||||
|  | 					StdErr:          stderr, | ||||||
|  | 					StdCombined:     stdcombined, | ||||||
|  | 					ExitCode:        exiterr.ExitCode(), | ||||||
|  | 					CommandTimedOut: false, | ||||||
|  | 				}, nil | ||||||
|  | 			} else if err != nil { | ||||||
|  | 				return CommandResult{}, err | ||||||
|  | 			} else { | ||||||
|  | 				return CommandResult{ | ||||||
|  | 					StdOut:          stdout, | ||||||
|  | 					StdErr:          stderr, | ||||||
|  | 					StdCombined:     stdcombined, | ||||||
|  | 					ExitCode:        0, | ||||||
|  | 					CommandTimedOut: false, | ||||||
|  | 				}, nil | ||||||
|  | 			} | ||||||
|  |  | ||||||
|  | 		} | ||||||
|  | 	} | ||||||
|  | } | ||||||
							
								
								
									
										172
									
								
								confext/confParser.go
									
									
									
									
									
										Normal file
									
								
							
							
						
						
									
										172
									
								
								confext/confParser.go
									
									
									
									
									
										Normal file
									
								
							| @@ -0,0 +1,172 @@ | |||||||
|  | package confext | ||||||
|  |  | ||||||
|  | import ( | ||||||
|  | 	"errors" | ||||||
|  | 	"fmt" | ||||||
|  | 	"gogs.mikescher.com/BlackForestBytes/goext/timeext" | ||||||
|  | 	"math/bits" | ||||||
|  | 	"os" | ||||||
|  | 	"reflect" | ||||||
|  | 	"strconv" | ||||||
|  | 	"time" | ||||||
|  | ) | ||||||
|  |  | ||||||
|  | // ApplyEnvOverrides overrides field values from environment variables | ||||||
|  | // | ||||||
|  | // fields must be tagged with `env:"env_key"` | ||||||
|  | // | ||||||
|  | // only works on exported fields | ||||||
|  | // | ||||||
|  | // fields without an env tag are ignored | ||||||
|  | // fields with an `env:"-"` tag are ignore | ||||||
|  | // | ||||||
|  | // sub-structs are recursively parsed (if they have an env tag) and the env-variable keys are delimited by the delim parameter | ||||||
|  | // sub-structs with `env:""` are also parsed, but the delimited is skipped (they are handled as if they were one level higher) | ||||||
|  | func ApplyEnvOverrides[T any](c *T, delim string) error { | ||||||
|  | 	rval := reflect.ValueOf(c).Elem() | ||||||
|  |  | ||||||
|  | 	return processEnvOverrides(rval, delim, "") | ||||||
|  | } | ||||||
|  |  | ||||||
|  | func processEnvOverrides(rval reflect.Value, delim string, prefix string) error { | ||||||
|  | 	rtyp := rval.Type() | ||||||
|  |  | ||||||
|  | 	for i := 0; i < rtyp.NumField(); i++ { | ||||||
|  |  | ||||||
|  | 		rsfield := rtyp.Field(i) | ||||||
|  | 		rvfield := rval.Field(i) | ||||||
|  |  | ||||||
|  | 		if !rsfield.IsExported() { | ||||||
|  | 			continue | ||||||
|  | 		} | ||||||
|  |  | ||||||
|  | 		if rvfield.Kind() == reflect.Struct { | ||||||
|  |  | ||||||
|  | 			envkey, found := rsfield.Tag.Lookup("env") | ||||||
|  | 			if !found || envkey == "-" { | ||||||
|  | 				continue | ||||||
|  | 			} | ||||||
|  |  | ||||||
|  | 			subPrefix := prefix | ||||||
|  | 			if envkey != "" { | ||||||
|  | 				subPrefix = subPrefix + envkey + delim | ||||||
|  | 			} | ||||||
|  |  | ||||||
|  | 			err := processEnvOverrides(rvfield, delim, subPrefix) | ||||||
|  | 			if err != nil { | ||||||
|  | 				return err | ||||||
|  | 			} | ||||||
|  | 		} | ||||||
|  |  | ||||||
|  | 		envkey := rsfield.Tag.Get("env") | ||||||
|  | 		if envkey == "" || envkey == "-" { | ||||||
|  | 			continue | ||||||
|  | 		} | ||||||
|  |  | ||||||
|  | 		fullEnvKey := prefix + envkey | ||||||
|  |  | ||||||
|  | 		envval, efound := os.LookupEnv(fullEnvKey) | ||||||
|  | 		if !efound { | ||||||
|  | 			continue | ||||||
|  | 		} | ||||||
|  |  | ||||||
|  | 		if rvfield.Type() == reflect.TypeOf("") { | ||||||
|  |  | ||||||
|  | 			rvfield.Set(reflect.ValueOf(envval)) | ||||||
|  |  | ||||||
|  | 			fmt.Printf("[CONF] Overwrite config '%s' with '%s'\n", fullEnvKey, envval) | ||||||
|  |  | ||||||
|  | 		} else if rvfield.Type() == reflect.TypeOf(int(0)) { | ||||||
|  |  | ||||||
|  | 			envint, err := strconv.ParseInt(envval, 10, bits.UintSize) | ||||||
|  | 			if err != nil { | ||||||
|  | 				return errors.New(fmt.Sprintf("Failed to parse env-config variable '%s' to int (value := '%s')", fullEnvKey, envval)) | ||||||
|  | 			} | ||||||
|  |  | ||||||
|  | 			rvfield.Set(reflect.ValueOf(int(envint))) | ||||||
|  |  | ||||||
|  | 			fmt.Printf("[CONF] Overwrite config '%s' with '%s'\n", fullEnvKey, envval) | ||||||
|  |  | ||||||
|  | 		} else if rvfield.Type() == reflect.TypeOf(int64(0)) { | ||||||
|  |  | ||||||
|  | 			envint, err := strconv.ParseInt(envval, 10, 64) | ||||||
|  | 			if err != nil { | ||||||
|  | 				return errors.New(fmt.Sprintf("Failed to parse env-config variable '%s' to int64 (value := '%s')", fullEnvKey, envval)) | ||||||
|  | 			} | ||||||
|  |  | ||||||
|  | 			rvfield.Set(reflect.ValueOf(int64(envint))) | ||||||
|  |  | ||||||
|  | 			fmt.Printf("[CONF] Overwrite config '%s' with '%s'\n", fullEnvKey, envval) | ||||||
|  |  | ||||||
|  | 		} else if rvfield.Type() == reflect.TypeOf(int32(0)) { | ||||||
|  |  | ||||||
|  | 			envint, err := strconv.ParseInt(envval, 10, 32) | ||||||
|  | 			if err != nil { | ||||||
|  | 				return errors.New(fmt.Sprintf("Failed to parse env-config variable '%s' to int32 (value := '%s')", fullEnvKey, envval)) | ||||||
|  | 			} | ||||||
|  |  | ||||||
|  | 			rvfield.Set(reflect.ValueOf(int32(envint))) | ||||||
|  |  | ||||||
|  | 			fmt.Printf("[CONF] Overwrite config '%s' with '%s'\n", fullEnvKey, envval) | ||||||
|  |  | ||||||
|  | 		} else if rvfield.Type() == reflect.TypeOf(int8(0)) { | ||||||
|  |  | ||||||
|  | 			envint, err := strconv.ParseInt(envval, 10, 8) | ||||||
|  | 			if err != nil { | ||||||
|  | 				return errors.New(fmt.Sprintf("Failed to parse env-config variable '%s' to int32 (value := '%s')", fullEnvKey, envval)) | ||||||
|  | 			} | ||||||
|  |  | ||||||
|  | 			rvfield.Set(reflect.ValueOf(int8(envint))) | ||||||
|  |  | ||||||
|  | 			fmt.Printf("[CONF] Overwrite config '%s' with '%s'\n", fullEnvKey, envval) | ||||||
|  |  | ||||||
|  | 		} else if rvfield.Type() == reflect.TypeOf(time.Duration(0)) { | ||||||
|  |  | ||||||
|  | 			dur, err := timeext.ParseDurationShortString(envval) | ||||||
|  | 			if err != nil { | ||||||
|  | 				return errors.New(fmt.Sprintf("Failed to parse env-config variable '%s' to duration (value := '%s')", fullEnvKey, envval)) | ||||||
|  | 			} | ||||||
|  |  | ||||||
|  | 			rvfield.Set(reflect.ValueOf(dur)) | ||||||
|  |  | ||||||
|  | 			fmt.Printf("[CONF] Overwrite config '%s' with '%s'\n", fullEnvKey, dur.String()) | ||||||
|  |  | ||||||
|  | 		} else if rvfield.Type() == reflect.TypeOf(time.UnixMilli(0)) { | ||||||
|  |  | ||||||
|  | 			tim, err := time.Parse(time.RFC3339Nano, envval) | ||||||
|  | 			if err != nil { | ||||||
|  | 				return errors.New(fmt.Sprintf("Failed to parse env-config variable '%s' to time.time (value := '%s')", fullEnvKey, envval)) | ||||||
|  | 			} | ||||||
|  |  | ||||||
|  | 			rvfield.Set(reflect.ValueOf(tim)) | ||||||
|  |  | ||||||
|  | 			fmt.Printf("[CONF] Overwrite config '%s' with '%s'\n", fullEnvKey, tim.String()) | ||||||
|  |  | ||||||
|  | 		} else if rvfield.Type().ConvertibleTo(reflect.TypeOf(int(0))) { | ||||||
|  |  | ||||||
|  | 			envint, err := strconv.ParseInt(envval, 10, 8) | ||||||
|  | 			if err != nil { | ||||||
|  | 				return errors.New(fmt.Sprintf("Failed to parse env-config variable '%s' to <%s, ,int> (value := '%s')", rvfield.Type().Name(), fullEnvKey, envval)) | ||||||
|  | 			} | ||||||
|  |  | ||||||
|  | 			envcvl := reflect.ValueOf(envint).Convert(rvfield.Type()) | ||||||
|  |  | ||||||
|  | 			rvfield.Set(envcvl) | ||||||
|  |  | ||||||
|  | 			fmt.Printf("[CONF] Overwrite config '%s' with '%v'\n", fullEnvKey, envcvl.Interface()) | ||||||
|  |  | ||||||
|  | 		} else if rvfield.Type().ConvertibleTo(reflect.TypeOf("")) { | ||||||
|  |  | ||||||
|  | 			envcvl := reflect.ValueOf(envval).Convert(rvfield.Type()) | ||||||
|  |  | ||||||
|  | 			rvfield.Set(envcvl) | ||||||
|  |  | ||||||
|  | 			fmt.Printf("[CONF] Overwrite config '%s' with '%v'\n", fullEnvKey, envcvl.Interface()) | ||||||
|  |  | ||||||
|  | 		} else { | ||||||
|  | 			return errors.New(fmt.Sprintf("Unknown kind/type in config: [ %s | %s ]", rvfield.Kind().String(), rvfield.Type().String())) | ||||||
|  | 		} | ||||||
|  | 	} | ||||||
|  |  | ||||||
|  | 	return nil | ||||||
|  | } | ||||||
							
								
								
									
										220
									
								
								confext/confParser_test.go
									
									
									
									
									
										Normal file
									
								
							
							
						
						
									
										220
									
								
								confext/confParser_test.go
									
									
									
									
									
										Normal file
									
								
							| @@ -0,0 +1,220 @@ | |||||||
|  | package confext | ||||||
|  |  | ||||||
|  | import ( | ||||||
|  | 	"gogs.mikescher.com/BlackForestBytes/goext/timeext" | ||||||
|  | 	"testing" | ||||||
|  | 	"time" | ||||||
|  | ) | ||||||
|  |  | ||||||
|  | func TestApplyEnvOverridesNoop(t *testing.T) { | ||||||
|  |  | ||||||
|  | 	type aliasint int | ||||||
|  | 	type aliasstring string | ||||||
|  |  | ||||||
|  | 	type testdata struct { | ||||||
|  | 		V1 int           `env:"TEST_V1"` | ||||||
|  | 		VX string        `` | ||||||
|  | 		V2 string        `env:"TEST_V2"` | ||||||
|  | 		V3 int8          `env:"TEST_V3"` | ||||||
|  | 		V4 int32         `env:"TEST_V4"` | ||||||
|  | 		V5 int64         `env:"TEST_V5"` | ||||||
|  | 		V6 aliasint      `env:"TEST_V6"` | ||||||
|  | 		VY aliasint      `` | ||||||
|  | 		V7 aliasstring   `env:"TEST_V7"` | ||||||
|  | 		V8 time.Duration `env:"TEST_V8"` | ||||||
|  | 		V9 time.Time     `env:"TEST_V9"` | ||||||
|  | 	} | ||||||
|  |  | ||||||
|  | 	input := testdata{ | ||||||
|  | 		V1: 1, | ||||||
|  | 		VX: "X", | ||||||
|  | 		V2: "2", | ||||||
|  | 		V3: 3, | ||||||
|  | 		V4: 4, | ||||||
|  | 		V5: 5, | ||||||
|  | 		V6: 6, | ||||||
|  | 		VY: 99, | ||||||
|  | 		V7: "7", | ||||||
|  | 		V8: 9, | ||||||
|  | 		V9: time.Unix(1671102873, 0), | ||||||
|  | 	} | ||||||
|  |  | ||||||
|  | 	output := input | ||||||
|  |  | ||||||
|  | 	err := ApplyEnvOverrides(&output, ".") | ||||||
|  | 	if err != nil { | ||||||
|  | 		t.Errorf("%v", err) | ||||||
|  | 		t.FailNow() | ||||||
|  | 	} | ||||||
|  |  | ||||||
|  | 	assertEqual(t, input, output) | ||||||
|  | } | ||||||
|  |  | ||||||
|  | func TestApplyEnvOverridesSimple(t *testing.T) { | ||||||
|  |  | ||||||
|  | 	type aliasint int | ||||||
|  | 	type aliasstring string | ||||||
|  |  | ||||||
|  | 	type testdata struct { | ||||||
|  | 		V1 int           `env:"TEST_V1"` | ||||||
|  | 		VX string        `` | ||||||
|  | 		V2 string        `env:"TEST_V2"` | ||||||
|  | 		V3 int8          `env:"TEST_V3"` | ||||||
|  | 		V4 int32         `env:"TEST_V4"` | ||||||
|  | 		V5 int64         `env:"TEST_V5"` | ||||||
|  | 		V6 aliasint      `env:"TEST_V6"` | ||||||
|  | 		VY aliasint      `` | ||||||
|  | 		V7 aliasstring   `env:"TEST_V7"` | ||||||
|  | 		V8 time.Duration `env:"TEST_V8"` | ||||||
|  | 		V9 time.Time     `env:"TEST_V9"` | ||||||
|  | 	} | ||||||
|  |  | ||||||
|  | 	data := testdata{ | ||||||
|  | 		V1: 1, | ||||||
|  | 		VX: "X", | ||||||
|  | 		V2: "2", | ||||||
|  | 		V3: 3, | ||||||
|  | 		V4: 4, | ||||||
|  | 		V5: 5, | ||||||
|  | 		V6: 6, | ||||||
|  | 		VY: 99, | ||||||
|  | 		V7: "7", | ||||||
|  | 		V8: 9, | ||||||
|  | 		V9: time.Unix(1671102873, 0), | ||||||
|  | 	} | ||||||
|  |  | ||||||
|  | 	t.Setenv("TEST_V1", "846") | ||||||
|  | 	t.Setenv("TEST_V2", "hello_world") | ||||||
|  | 	t.Setenv("TEST_V3", "6") | ||||||
|  | 	t.Setenv("TEST_V4", "333") | ||||||
|  | 	t.Setenv("TEST_V5", "-937") | ||||||
|  | 	t.Setenv("TEST_V6", "070") | ||||||
|  | 	t.Setenv("TEST_V7", "AAAAAA") | ||||||
|  | 	t.Setenv("TEST_V8", "1min4s") | ||||||
|  | 	t.Setenv("TEST_V9", "2009-11-10T23:00:00Z") | ||||||
|  |  | ||||||
|  | 	err := ApplyEnvOverrides(&data, ".") | ||||||
|  | 	if err != nil { | ||||||
|  | 		t.Errorf("%v", err) | ||||||
|  | 		t.FailNow() | ||||||
|  | 	} | ||||||
|  |  | ||||||
|  | 	assertEqual(t, data.V1, 846) | ||||||
|  | 	assertEqual(t, data.V2, "hello_world") | ||||||
|  | 	assertEqual(t, data.V3, 6) | ||||||
|  | 	assertEqual(t, data.V4, 333) | ||||||
|  | 	assertEqual(t, data.V5, -937) | ||||||
|  | 	assertEqual(t, data.V6, 70) | ||||||
|  | 	assertEqual(t, data.V7, "AAAAAA") | ||||||
|  | 	assertEqual(t, data.V8, time.Second*64) | ||||||
|  | 	assertEqual(t, data.V9, time.Unix(1257894000, 0).UTC()) | ||||||
|  | } | ||||||
|  |  | ||||||
|  | func TestApplyEnvOverridesRecursive(t *testing.T) { | ||||||
|  |  | ||||||
|  | 	type subdata struct { | ||||||
|  | 		V1 int           `env:"SUB_V1"` | ||||||
|  | 		VX string        `` | ||||||
|  | 		V2 string        `env:"SUB_V2"` | ||||||
|  | 		V8 time.Duration `env:"SUB_V3"` | ||||||
|  | 		V9 time.Time     `env:"SUB_V4"` | ||||||
|  | 	} | ||||||
|  |  | ||||||
|  | 	type testdata struct { | ||||||
|  | 		V1   int     `env:"TEST_V1"` | ||||||
|  | 		VX   string  `` | ||||||
|  | 		Sub1 subdata `` | ||||||
|  | 		Sub2 subdata `env:"TEST_V2"` | ||||||
|  | 		Sub3 subdata `env:"TEST_V3"` | ||||||
|  | 		Sub4 subdata `env:""` | ||||||
|  | 		V5   string  `env:"-"` | ||||||
|  | 	} | ||||||
|  |  | ||||||
|  | 	data := testdata{ | ||||||
|  | 		V1: 1, | ||||||
|  | 		VX: "2", | ||||||
|  | 		V5: "no", | ||||||
|  | 		Sub1: subdata{ | ||||||
|  | 			V1: 3, | ||||||
|  | 			VX: "4", | ||||||
|  | 			V2: "5", | ||||||
|  | 			V8: 6 * time.Second, | ||||||
|  | 			V9: time.Date(2000, 1, 7, 1, 1, 1, 0, time.UTC), | ||||||
|  | 		}, | ||||||
|  | 		Sub2: subdata{ | ||||||
|  | 			V1: 8, | ||||||
|  | 			VX: "9", | ||||||
|  | 			V2: "10", | ||||||
|  | 			V8: 11 * time.Second, | ||||||
|  | 			V9: time.Date(2000, 1, 12, 1, 1, 1, 0, timeext.TimezoneBerlin), | ||||||
|  | 		}, | ||||||
|  | 		Sub3: subdata{ | ||||||
|  | 			V1: 13, | ||||||
|  | 			VX: "14", | ||||||
|  | 			V2: "15", | ||||||
|  | 			V8: 16 * time.Second, | ||||||
|  | 			V9: time.Date(2000, 1, 17, 1, 1, 1, 0, timeext.TimezoneBerlin), | ||||||
|  | 		}, | ||||||
|  | 		Sub4: subdata{ | ||||||
|  | 			V1: 18, | ||||||
|  | 			VX: "19", | ||||||
|  | 			V2: "20", | ||||||
|  | 			V8: 21 * time.Second, | ||||||
|  | 			V9: time.Date(2000, 1, 22, 1, 1, 1, 0, timeext.TimezoneBerlin), | ||||||
|  | 		}, | ||||||
|  | 	} | ||||||
|  |  | ||||||
|  | 	t.Setenv("TEST_V1", "999") | ||||||
|  | 	t.Setenv("-", "yes") | ||||||
|  |  | ||||||
|  | 	t.Setenv("TEST_V2_SUB_V1", "846") | ||||||
|  | 	t.Setenv("TEST_V2_SUB_V2", "222_hello_world") | ||||||
|  | 	t.Setenv("TEST_V2_SUB_V3", "1min4s") | ||||||
|  | 	t.Setenv("TEST_V2_SUB_V4", "2009-11-10T23:00:00Z") | ||||||
|  |  | ||||||
|  | 	t.Setenv("TEST_V3_SUB_V1", "33846") | ||||||
|  | 	t.Setenv("TEST_V3_SUB_V2", "33_hello_world") | ||||||
|  | 	t.Setenv("TEST_V3_SUB_V3", "33min4s") | ||||||
|  | 	t.Setenv("TEST_V3_SUB_V4", "2033-11-10T23:00:00Z") | ||||||
|  |  | ||||||
|  | 	t.Setenv("SUB_V1", "11") | ||||||
|  | 	t.Setenv("SUB_V2", "22") | ||||||
|  | 	t.Setenv("SUB_V3", "33min") | ||||||
|  | 	t.Setenv("SUB_V4", "2044-01-01T00:00:00Z") | ||||||
|  |  | ||||||
|  | 	err := ApplyEnvOverrides(&data, "_") | ||||||
|  | 	if err != nil { | ||||||
|  | 		t.Errorf("%v", err) | ||||||
|  | 		t.FailNow() | ||||||
|  | 	} | ||||||
|  |  | ||||||
|  | 	assertEqual(t, data.V1, 999) | ||||||
|  | 	assertEqual(t, data.VX, "2") | ||||||
|  | 	assertEqual(t, data.V5, "no") | ||||||
|  | 	assertEqual(t, data.Sub1.V1, 3) | ||||||
|  | 	assertEqual(t, data.Sub1.VX, "4") | ||||||
|  | 	assertEqual(t, data.Sub1.V2, "5") | ||||||
|  | 	assertEqual(t, data.Sub1.V8, time.Second*6) | ||||||
|  | 	assertEqual(t, data.Sub1.V9, time.Unix(947206861, 0).UTC()) | ||||||
|  | 	assertEqual(t, data.Sub2.V1, 846) | ||||||
|  | 	assertEqual(t, data.Sub2.VX, "9") | ||||||
|  | 	assertEqual(t, data.Sub2.V2, "222_hello_world") | ||||||
|  | 	assertEqual(t, data.Sub2.V8, time.Second*64) | ||||||
|  | 	assertEqual(t, data.Sub2.V9, time.Unix(1257894000, 0).UTC()) | ||||||
|  | 	assertEqual(t, data.Sub3.V1, 33846) | ||||||
|  | 	assertEqual(t, data.Sub3.VX, "14") | ||||||
|  | 	assertEqual(t, data.Sub3.V2, "33_hello_world") | ||||||
|  | 	assertEqual(t, data.Sub3.V8, time.Second*1984) | ||||||
|  | 	assertEqual(t, data.Sub3.V9, time.Unix(2015276400, 0).UTC()) | ||||||
|  | 	assertEqual(t, data.Sub4.V1, 11) | ||||||
|  | 	assertEqual(t, data.Sub4.VX, "19") | ||||||
|  | 	assertEqual(t, data.Sub4.V2, "22") | ||||||
|  | 	assertEqual(t, data.Sub4.V8, time.Second*1980) | ||||||
|  | 	assertEqual(t, data.Sub4.V9, time.Unix(2335219200, 0).UTC()) | ||||||
|  | } | ||||||
|  |  | ||||||
|  | func assertEqual[T comparable](t *testing.T, actual T, expected T) { | ||||||
|  | 	if actual != expected { | ||||||
|  | 		t.Errorf("values differ: Actual: '%v', Expected: '%v'", actual, expected) | ||||||
|  | 	} | ||||||
|  | } | ||||||
							
								
								
									
										365
									
								
								cryptext/passHash.go
									
									
									
									
									
										Normal file
									
								
							
							
						
						
									
										365
									
								
								cryptext/passHash.go
									
									
									
									
									
										Normal file
									
								
							| @@ -0,0 +1,365 @@ | |||||||
|  | package cryptext | ||||||
|  |  | ||||||
|  | import ( | ||||||
|  | 	"crypto/rand" | ||||||
|  | 	"crypto/sha256" | ||||||
|  | 	"encoding/base64" | ||||||
|  | 	"encoding/hex" | ||||||
|  | 	"errors" | ||||||
|  | 	"fmt" | ||||||
|  | 	"gogs.mikescher.com/BlackForestBytes/goext/langext" | ||||||
|  | 	"gogs.mikescher.com/BlackForestBytes/goext/totpext" | ||||||
|  | 	"golang.org/x/crypto/bcrypt" | ||||||
|  | 	"strconv" | ||||||
|  | 	"strings" | ||||||
|  | ) | ||||||
|  |  | ||||||
|  | const LatestPassHashVersion = 4 | ||||||
|  |  | ||||||
|  | // PassHash | ||||||
|  | // - [v0]: plaintext password ( `0|...` ) | ||||||
|  | // - [v1]: sha256(plaintext) | ||||||
|  | // - [v2]: seed | sha256<seed>(plaintext) | ||||||
|  | // - [v3]: seed | sha256<seed>(plaintext) | [hex(totp)] | ||||||
|  | // - [v4]: bcrypt(plaintext) | [hex(totp)] | ||||||
|  | type PassHash string | ||||||
|  |  | ||||||
|  | func (ph PassHash) Valid() bool { | ||||||
|  | 	_, _, _, _, _, valid := ph.Data() | ||||||
|  | 	return valid | ||||||
|  | } | ||||||
|  |  | ||||||
|  | func (ph PassHash) HasTOTP() bool { | ||||||
|  | 	_, _, _, otp, _, _ := ph.Data() | ||||||
|  | 	return otp | ||||||
|  | } | ||||||
|  |  | ||||||
|  | func (ph PassHash) Data() (_version int, _seed []byte, _payload []byte, _totp bool, _totpsecret []byte, _valid bool) { | ||||||
|  |  | ||||||
|  | 	split := strings.Split(string(ph), "|") | ||||||
|  | 	if len(split) == 0 { | ||||||
|  | 		return -1, nil, nil, false, nil, false | ||||||
|  | 	} | ||||||
|  |  | ||||||
|  | 	version, err := strconv.ParseInt(split[0], 10, 32) | ||||||
|  | 	if err != nil { | ||||||
|  | 		return -1, nil, nil, false, nil, false | ||||||
|  | 	} | ||||||
|  |  | ||||||
|  | 	if version == 0 { | ||||||
|  | 		if len(split) != 2 { | ||||||
|  | 			return -1, nil, nil, false, nil, false | ||||||
|  | 		} | ||||||
|  | 		return int(version), nil, []byte(split[1]), false, nil, true | ||||||
|  | 	} | ||||||
|  |  | ||||||
|  | 	if version == 1 { | ||||||
|  | 		if len(split) != 2 { | ||||||
|  | 			return -1, nil, nil, false, nil, false | ||||||
|  | 		} | ||||||
|  | 		payload, err := base64.RawStdEncoding.DecodeString(split[1]) | ||||||
|  | 		if err != nil { | ||||||
|  | 			return -1, nil, nil, false, nil, false | ||||||
|  | 		} | ||||||
|  | 		return int(version), nil, payload, false, nil, true | ||||||
|  | 	} | ||||||
|  |  | ||||||
|  | 	// | ||||||
|  | 	if version == 2 { | ||||||
|  | 		if len(split) != 3 { | ||||||
|  | 			return -1, nil, nil, false, nil, false | ||||||
|  | 		} | ||||||
|  | 		seed, err := base64.RawStdEncoding.DecodeString(split[1]) | ||||||
|  | 		if err != nil { | ||||||
|  | 			return -1, nil, nil, false, nil, false | ||||||
|  | 		} | ||||||
|  | 		payload, err := base64.RawStdEncoding.DecodeString(split[2]) | ||||||
|  | 		if err != nil { | ||||||
|  | 			return -1, nil, nil, false, nil, false | ||||||
|  | 		} | ||||||
|  | 		return int(version), seed, payload, false, nil, true | ||||||
|  | 	} | ||||||
|  |  | ||||||
|  | 	if version == 3 { | ||||||
|  | 		if len(split) != 4 { | ||||||
|  | 			return -1, nil, nil, false, nil, false | ||||||
|  | 		} | ||||||
|  | 		seed, err := base64.RawStdEncoding.DecodeString(split[1]) | ||||||
|  | 		if err != nil { | ||||||
|  | 			return -1, nil, nil, false, nil, false | ||||||
|  | 		} | ||||||
|  | 		payload, err := base64.RawStdEncoding.DecodeString(split[2]) | ||||||
|  | 		if err != nil { | ||||||
|  | 			return -1, nil, nil, false, nil, false | ||||||
|  | 		} | ||||||
|  | 		totp := false | ||||||
|  | 		totpsecret := make([]byte, 0) | ||||||
|  | 		if split[3] != "0" { | ||||||
|  | 			totpsecret, err = hex.DecodeString(split[3]) | ||||||
|  | 			totp = true | ||||||
|  | 		} | ||||||
|  | 		return int(version), seed, payload, totp, totpsecret, true | ||||||
|  | 	} | ||||||
|  |  | ||||||
|  | 	if version == 4 { | ||||||
|  | 		if len(split) != 3 { | ||||||
|  | 			return -1, nil, nil, false, nil, false | ||||||
|  | 		} | ||||||
|  | 		payload := []byte(split[1]) | ||||||
|  | 		totp := false | ||||||
|  | 		totpsecret := make([]byte, 0) | ||||||
|  | 		if split[2] != "0" { | ||||||
|  | 			totpsecret, err = hex.DecodeString(split[3]) | ||||||
|  | 			totp = true | ||||||
|  | 		} | ||||||
|  | 		return int(version), nil, payload, totp, totpsecret, true | ||||||
|  | 	} | ||||||
|  |  | ||||||
|  | 	return -1, nil, nil, false, nil, false | ||||||
|  | } | ||||||
|  |  | ||||||
|  | func (ph PassHash) Verify(plainpass string, totp *string) bool { | ||||||
|  | 	version, seed, payload, hastotp, totpsecret, valid := ph.Data() | ||||||
|  | 	if !valid { | ||||||
|  | 		return false | ||||||
|  | 	} | ||||||
|  |  | ||||||
|  | 	if hastotp && totp == nil { | ||||||
|  | 		return false | ||||||
|  | 	} | ||||||
|  |  | ||||||
|  | 	if version == 0 { | ||||||
|  | 		return langext.ArrEqualsExact([]byte(plainpass), payload) | ||||||
|  | 	} | ||||||
|  |  | ||||||
|  | 	if version == 1 { | ||||||
|  | 		return langext.ArrEqualsExact(hash256(plainpass), payload) | ||||||
|  | 	} | ||||||
|  |  | ||||||
|  | 	if version == 2 { | ||||||
|  | 		return langext.ArrEqualsExact(hash256Seeded(plainpass, seed), payload) | ||||||
|  | 	} | ||||||
|  |  | ||||||
|  | 	if version == 3 { | ||||||
|  | 		if !hastotp { | ||||||
|  | 			return langext.ArrEqualsExact(hash256Seeded(plainpass, seed), payload) | ||||||
|  | 		} else { | ||||||
|  | 			return langext.ArrEqualsExact(hash256Seeded(plainpass, seed), payload) && totpext.Validate(totpsecret, *totp) | ||||||
|  | 		} | ||||||
|  | 	} | ||||||
|  |  | ||||||
|  | 	if version == 4 { | ||||||
|  | 		if !hastotp { | ||||||
|  | 			return bcrypt.CompareHashAndPassword(payload, []byte(plainpass)) == nil | ||||||
|  | 		} else { | ||||||
|  | 			return bcrypt.CompareHashAndPassword(payload, []byte(plainpass)) == nil && totpext.Validate(totpsecret, *totp) | ||||||
|  | 		} | ||||||
|  | 	} | ||||||
|  |  | ||||||
|  | 	return false | ||||||
|  | } | ||||||
|  |  | ||||||
|  | func (ph PassHash) NeedsPasswordUpgrade() bool { | ||||||
|  | 	version, _, _, _, _, valid := ph.Data() | ||||||
|  | 	return valid && version < LatestPassHashVersion | ||||||
|  | } | ||||||
|  |  | ||||||
|  | func (ph PassHash) Upgrade(plainpass string) (PassHash, error) { | ||||||
|  | 	version, _, _, hastotp, totpsecret, valid := ph.Data() | ||||||
|  | 	if !valid { | ||||||
|  | 		return "", errors.New("invalid password") | ||||||
|  | 	} | ||||||
|  | 	if version == LatestPassHashVersion { | ||||||
|  | 		return ph, nil | ||||||
|  | 	} | ||||||
|  | 	if hastotp { | ||||||
|  | 		return HashPassword(plainpass, totpsecret) | ||||||
|  | 	} else { | ||||||
|  | 		return HashPassword(plainpass, nil) | ||||||
|  | 	} | ||||||
|  | } | ||||||
|  |  | ||||||
|  | func (ph PassHash) ClearTOTP() (PassHash, error) { | ||||||
|  | 	version, _, _, _, _, valid := ph.Data() | ||||||
|  | 	if !valid { | ||||||
|  | 		return "", errors.New("invalid PassHash") | ||||||
|  | 	} | ||||||
|  |  | ||||||
|  | 	if version == 0 { | ||||||
|  | 		return ph, nil | ||||||
|  | 	} | ||||||
|  |  | ||||||
|  | 	if version == 1 { | ||||||
|  | 		return ph, nil | ||||||
|  | 	} | ||||||
|  |  | ||||||
|  | 	if version == 2 { | ||||||
|  | 		return ph, nil | ||||||
|  | 	} | ||||||
|  |  | ||||||
|  | 	if version == 3 { | ||||||
|  | 		split := strings.Split(string(ph), "|") | ||||||
|  | 		split[3] = "0" | ||||||
|  | 		return PassHash(strings.Join(split, "|")), nil | ||||||
|  | 	} | ||||||
|  |  | ||||||
|  | 	if version == 4 { | ||||||
|  | 		split := strings.Split(string(ph), "|") | ||||||
|  | 		split[2] = "0" | ||||||
|  | 		return PassHash(strings.Join(split, "|")), nil | ||||||
|  | 	} | ||||||
|  |  | ||||||
|  | 	return "", errors.New("unknown version") | ||||||
|  | } | ||||||
|  |  | ||||||
|  | func (ph PassHash) WithTOTP(totpSecret []byte) (PassHash, error) { | ||||||
|  | 	version, _, _, _, _, valid := ph.Data() | ||||||
|  | 	if !valid { | ||||||
|  | 		return "", errors.New("invalid PassHash") | ||||||
|  | 	} | ||||||
|  |  | ||||||
|  | 	if version == 0 { | ||||||
|  | 		return "", errors.New("version does not support totp, needs upgrade") | ||||||
|  | 	} | ||||||
|  |  | ||||||
|  | 	if version == 1 { | ||||||
|  | 		return "", errors.New("version does not support totp, needs upgrade") | ||||||
|  | 	} | ||||||
|  |  | ||||||
|  | 	if version == 2 { | ||||||
|  | 		return "", errors.New("version does not support totp, needs upgrade") | ||||||
|  | 	} | ||||||
|  |  | ||||||
|  | 	if version == 3 { | ||||||
|  | 		split := strings.Split(string(ph), "|") | ||||||
|  | 		split[3] = hex.EncodeToString(totpSecret) | ||||||
|  | 		return PassHash(strings.Join(split, "|")), nil | ||||||
|  | 	} | ||||||
|  |  | ||||||
|  | 	if version == 4 { | ||||||
|  | 		split := strings.Split(string(ph), "|") | ||||||
|  | 		split[2] = hex.EncodeToString(totpSecret) | ||||||
|  | 		return PassHash(strings.Join(split, "|")), nil | ||||||
|  | 	} | ||||||
|  |  | ||||||
|  | 	return "", errors.New("unknown version") | ||||||
|  | } | ||||||
|  |  | ||||||
|  | func (ph PassHash) Change(newPlainPass string) (PassHash, error) { | ||||||
|  | 	version, _, _, hastotp, totpsecret, valid := ph.Data() | ||||||
|  | 	if !valid { | ||||||
|  | 		return "", errors.New("invalid PassHash") | ||||||
|  | 	} | ||||||
|  |  | ||||||
|  | 	if version == 0 { | ||||||
|  | 		return HashPasswordV0(newPlainPass) | ||||||
|  | 	} | ||||||
|  |  | ||||||
|  | 	if version == 1 { | ||||||
|  | 		return HashPasswordV1(newPlainPass) | ||||||
|  | 	} | ||||||
|  |  | ||||||
|  | 	if version == 2 { | ||||||
|  | 		return HashPasswordV2(newPlainPass) | ||||||
|  | 	} | ||||||
|  |  | ||||||
|  | 	if version == 3 { | ||||||
|  | 		return HashPasswordV3(newPlainPass, langext.Conditional(hastotp, totpsecret, nil)) | ||||||
|  | 	} | ||||||
|  |  | ||||||
|  | 	if version == 4 { | ||||||
|  | 		return HashPasswordV4(newPlainPass, langext.Conditional(hastotp, totpsecret, nil)) | ||||||
|  | 	} | ||||||
|  |  | ||||||
|  | 	return "", errors.New("unknown version") | ||||||
|  | } | ||||||
|  |  | ||||||
|  | func (ph PassHash) String() string { | ||||||
|  | 	return string(ph) | ||||||
|  | } | ||||||
|  |  | ||||||
|  | func HashPassword(plainpass string, totpSecret []byte) (PassHash, error) { | ||||||
|  | 	return HashPasswordV4(plainpass, totpSecret) | ||||||
|  | } | ||||||
|  |  | ||||||
|  | func HashPasswordV4(plainpass string, totpSecret []byte) (PassHash, error) { | ||||||
|  | 	var strtotp string | ||||||
|  |  | ||||||
|  | 	if totpSecret == nil { | ||||||
|  | 		strtotp = "0" | ||||||
|  | 	} else { | ||||||
|  | 		strtotp = hex.EncodeToString(totpSecret) | ||||||
|  | 	} | ||||||
|  |  | ||||||
|  | 	payload, err := bcrypt.GenerateFromPassword([]byte(plainpass), bcrypt.MinCost) | ||||||
|  | 	if err != nil { | ||||||
|  | 		return "", err | ||||||
|  | 	} | ||||||
|  |  | ||||||
|  | 	return PassHash(fmt.Sprintf("4|%s|%s", string(payload), strtotp)), nil | ||||||
|  | } | ||||||
|  |  | ||||||
|  | func HashPasswordV3(plainpass string, totpSecret []byte) (PassHash, error) { | ||||||
|  | 	var strtotp string | ||||||
|  |  | ||||||
|  | 	if totpSecret == nil { | ||||||
|  | 		strtotp = "0" | ||||||
|  | 	} else { | ||||||
|  | 		strtotp = hex.EncodeToString(totpSecret) | ||||||
|  | 	} | ||||||
|  |  | ||||||
|  | 	seed, err := newSeed() | ||||||
|  | 	if err != nil { | ||||||
|  | 		return "", err | ||||||
|  | 	} | ||||||
|  |  | ||||||
|  | 	checksum := hash256Seeded(plainpass, seed) | ||||||
|  |  | ||||||
|  | 	return PassHash(fmt.Sprintf("3|%s|%s|%s", | ||||||
|  | 		base64.RawStdEncoding.EncodeToString(seed), | ||||||
|  | 		base64.RawStdEncoding.EncodeToString(checksum), | ||||||
|  | 		strtotp)), nil | ||||||
|  | } | ||||||
|  |  | ||||||
|  | func HashPasswordV2(plainpass string) (PassHash, error) { | ||||||
|  | 	seed, err := newSeed() | ||||||
|  | 	if err != nil { | ||||||
|  | 		return "", err | ||||||
|  | 	} | ||||||
|  |  | ||||||
|  | 	checksum := hash256Seeded(plainpass, seed) | ||||||
|  |  | ||||||
|  | 	return PassHash(fmt.Sprintf("2|%s|%s", base64.RawStdEncoding.EncodeToString(seed), base64.RawStdEncoding.EncodeToString(checksum))), nil | ||||||
|  | } | ||||||
|  |  | ||||||
|  | func HashPasswordV1(plainpass string) (PassHash, error) { | ||||||
|  | 	return PassHash(fmt.Sprintf("1|%s", base64.RawStdEncoding.EncodeToString(hash256(plainpass)))), nil | ||||||
|  | } | ||||||
|  |  | ||||||
|  | func HashPasswordV0(plainpass string) (PassHash, error) { | ||||||
|  | 	return PassHash(fmt.Sprintf("0|%s", plainpass)), nil | ||||||
|  | } | ||||||
|  |  | ||||||
|  | func hash256(s string) []byte { | ||||||
|  | 	h := sha256.New() | ||||||
|  | 	h.Write([]byte(s)) | ||||||
|  | 	bs := h.Sum(nil) | ||||||
|  | 	return bs | ||||||
|  | } | ||||||
|  |  | ||||||
|  | func hash256Seeded(s string, seed []byte) []byte { | ||||||
|  | 	h := sha256.New() | ||||||
|  | 	h.Write(seed) | ||||||
|  | 	h.Write([]byte(s)) | ||||||
|  | 	bs := h.Sum(nil) | ||||||
|  | 	return bs | ||||||
|  | } | ||||||
|  |  | ||||||
|  | func newSeed() ([]byte, error) { | ||||||
|  | 	secret := make([]byte, 32) | ||||||
|  | 	_, err := rand.Read(secret) | ||||||
|  | 	if err != nil { | ||||||
|  | 		return nil, err | ||||||
|  | 	} | ||||||
|  | 	return secret, nil | ||||||
|  | } | ||||||
| @@ -1,56 +1,157 @@ | |||||||
| package dataext | package dataext | ||||||
|  |  | ||||||
| import "io" | import ( | ||||||
|  | 	"errors" | ||||||
|  | 	"io" | ||||||
|  | ) | ||||||
|  |  | ||||||
|  | type brcMode int | ||||||
|  |  | ||||||
|  | const ( | ||||||
|  | 	modeSourceReading  brcMode = 0 | ||||||
|  | 	modeSourceFinished brcMode = 1 | ||||||
|  | 	modeBufferReading  brcMode = 2 | ||||||
|  | 	modeBufferFinished brcMode = 3 | ||||||
|  | ) | ||||||
|  |  | ||||||
| type BufferedReadCloser interface { | type BufferedReadCloser interface { | ||||||
| 	io.ReadCloser | 	io.ReadCloser | ||||||
| 	BufferedAll() ([]byte, error) | 	BufferedAll() ([]byte, error) | ||||||
|  | 	Reset() error | ||||||
| } | } | ||||||
|  |  | ||||||
| type bufferedReadCloser struct { | type bufferedReadCloser struct { | ||||||
| 	buffer []byte | 	buffer []byte | ||||||
| 	inner  io.ReadCloser | 	inner  io.ReadCloser | ||||||
| 	finished bool | 	mode   brcMode | ||||||
| } | 	off    int | ||||||
|  |  | ||||||
| func (b *bufferedReadCloser) Read(p []byte) (int, error) { |  | ||||||
|  |  | ||||||
| 	n, err := b.inner.Read(p) |  | ||||||
| 	if n > 0 { |  | ||||||
| 		b.buffer = append(b.buffer, p[0:n]...) |  | ||||||
| 	} |  | ||||||
|  |  | ||||||
| 	if err == io.EOF { |  | ||||||
| 		b.finished = true |  | ||||||
| 	} |  | ||||||
|  |  | ||||||
| 	return n, err |  | ||||||
| } | } | ||||||
|  |  | ||||||
| func NewBufferedReadCloser(sub io.ReadCloser) BufferedReadCloser { | func NewBufferedReadCloser(sub io.ReadCloser) BufferedReadCloser { | ||||||
| 	return &bufferedReadCloser{ | 	return &bufferedReadCloser{ | ||||||
| 		buffer: make([]byte, 0, 1024), | 		buffer: make([]byte, 0, 1024), | ||||||
| 		inner:  sub, | 		inner:  sub, | ||||||
| 		finished: false, | 		mode:   modeSourceReading, | ||||||
|  | 		off:    0, | ||||||
|  | 	} | ||||||
|  | } | ||||||
|  |  | ||||||
|  | func (b *bufferedReadCloser) Read(p []byte) (int, error) { | ||||||
|  | 	switch b.mode { | ||||||
|  | 	case modeSourceReading: | ||||||
|  | 		n, err := b.inner.Read(p) | ||||||
|  | 		if n > 0 { | ||||||
|  | 			b.buffer = append(b.buffer, p[0:n]...) | ||||||
|  | 		} | ||||||
|  |  | ||||||
|  | 		if err == io.EOF { | ||||||
|  | 			b.mode = modeSourceFinished | ||||||
|  | 		} | ||||||
|  |  | ||||||
|  | 		return n, err | ||||||
|  |  | ||||||
|  | 	case modeSourceFinished: | ||||||
|  | 		return 0, io.EOF | ||||||
|  |  | ||||||
|  | 	case modeBufferReading: | ||||||
|  |  | ||||||
|  | 		if len(b.buffer) <= b.off { | ||||||
|  | 			b.mode = modeBufferFinished | ||||||
|  | 			if len(p) == 0 { | ||||||
|  | 				return 0, nil | ||||||
|  | 			} | ||||||
|  | 			return 0, io.EOF | ||||||
|  | 		} | ||||||
|  |  | ||||||
|  | 		n := copy(p, b.buffer[b.off:]) | ||||||
|  | 		b.off += n | ||||||
|  | 		return n, nil | ||||||
|  |  | ||||||
|  | 	case modeBufferFinished: | ||||||
|  | 		return 0, io.EOF | ||||||
|  |  | ||||||
|  | 	default: | ||||||
|  | 		return 0, errors.New("object in undefined status") | ||||||
| 	} | 	} | ||||||
| } | } | ||||||
|  |  | ||||||
| func (b *bufferedReadCloser) Close() error { | func (b *bufferedReadCloser) Close() error { | ||||||
| 	err := b.inner.Close() | 	switch b.mode { | ||||||
|  | 	case modeSourceReading: | ||||||
|  | 		_, err := b.BufferedAll() | ||||||
| 		if err != nil { | 		if err != nil { | ||||||
| 		b.finished = true |  | ||||||
| 	} |  | ||||||
| 			return err | 			return err | ||||||
| 		} | 		} | ||||||
|  | 		err = b.inner.Close() | ||||||
|  | 		if err != nil { | ||||||
|  | 			return err | ||||||
|  | 		} | ||||||
|  | 		b.mode = modeSourceFinished | ||||||
|  | 		return nil | ||||||
|  |  | ||||||
|  | 	case modeSourceFinished: | ||||||
|  | 		return nil | ||||||
|  |  | ||||||
|  | 	case modeBufferReading: | ||||||
|  | 		b.mode = modeBufferFinished | ||||||
|  | 		return nil | ||||||
|  |  | ||||||
|  | 	case modeBufferFinished: | ||||||
|  | 		return nil | ||||||
|  |  | ||||||
|  | 	default: | ||||||
|  | 		return errors.New("object in undefined status") | ||||||
|  | 	} | ||||||
|  |  | ||||||
|  | } | ||||||
|  |  | ||||||
| func (b *bufferedReadCloser) BufferedAll() ([]byte, error) { | func (b *bufferedReadCloser) BufferedAll() ([]byte, error) { | ||||||
|  | 	switch b.mode { | ||||||
|  | 	case modeSourceReading: | ||||||
| 		arr := make([]byte, 1024) | 		arr := make([]byte, 1024) | ||||||
| 	for !b.finished { | 		for b.mode == modeSourceReading { | ||||||
| 			_, err := b.Read(arr) | 			_, err := b.Read(arr) | ||||||
| 			if err != nil && err != io.EOF { | 			if err != nil && err != io.EOF { | ||||||
| 				return nil, err | 				return nil, err | ||||||
| 			} | 			} | ||||||
| 		} | 		} | ||||||
|  |  | ||||||
| 		return b.buffer, nil | 		return b.buffer, nil | ||||||
|  |  | ||||||
|  | 	case modeSourceFinished: | ||||||
|  | 		return b.buffer, nil | ||||||
|  |  | ||||||
|  | 	case modeBufferReading: | ||||||
|  | 		return b.buffer, nil | ||||||
|  |  | ||||||
|  | 	case modeBufferFinished: | ||||||
|  | 		return b.buffer, nil | ||||||
|  |  | ||||||
|  | 	default: | ||||||
|  | 		return nil, errors.New("object in undefined status") | ||||||
|  | 	} | ||||||
|  | } | ||||||
|  |  | ||||||
|  | func (b *bufferedReadCloser) Reset() error { | ||||||
|  | 	switch b.mode { | ||||||
|  | 	case modeSourceReading: | ||||||
|  | 		fallthrough | ||||||
|  | 	case modeSourceFinished: | ||||||
|  | 		err := b.Close() | ||||||
|  | 		if err != nil { | ||||||
|  | 			return err | ||||||
|  | 		} | ||||||
|  | 		b.mode = modeBufferReading | ||||||
|  | 		b.off = 0 | ||||||
|  | 		return nil | ||||||
|  |  | ||||||
|  | 	case modeBufferReading: | ||||||
|  | 		fallthrough | ||||||
|  | 	case modeBufferFinished: | ||||||
|  | 		b.mode = modeBufferReading | ||||||
|  | 		b.off = 0 | ||||||
|  | 		return nil | ||||||
|  |  | ||||||
|  | 	default: | ||||||
|  | 		return errors.New("object in undefined status") | ||||||
|  | 	} | ||||||
| } | } | ||||||
|   | |||||||
| @@ -19,40 +19,38 @@ import ( | |||||||
| // There are also a bunch of unit tests to ensure that the cache is always in a consistent state | // There are also a bunch of unit tests to ensure that the cache is always in a consistent state | ||||||
| // | // | ||||||
|  |  | ||||||
| type LRUData interface{} | type LRUMap[TKey comparable, TData any] struct { | ||||||
|  |  | ||||||
| type LRUMap struct { |  | ||||||
| 	maxsize int | 	maxsize int | ||||||
| 	lock    sync.Mutex | 	lock    sync.Mutex | ||||||
|  |  | ||||||
| 	cache map[string]*cacheNode | 	cache map[TKey]*cacheNode[TKey, TData] | ||||||
|  |  | ||||||
| 	lfuHead *cacheNode | 	lfuHead *cacheNode[TKey, TData] | ||||||
| 	lfuTail *cacheNode | 	lfuTail *cacheNode[TKey, TData] | ||||||
| } | } | ||||||
|  |  | ||||||
| type cacheNode struct { | type cacheNode[TKey comparable, TData any] struct { | ||||||
| 	key    string | 	key    TKey | ||||||
| 	data   LRUData | 	data   TData | ||||||
| 	parent *cacheNode | 	parent *cacheNode[TKey, TData] | ||||||
| 	child  *cacheNode | 	child  *cacheNode[TKey, TData] | ||||||
| } | } | ||||||
|  |  | ||||||
| func NewLRUMap(size int) *LRUMap { | func NewLRUMap[TKey comparable, TData any](size int) *LRUMap[TKey, TData] { | ||||||
| 	if size <= 2 && size != 0 { | 	if size <= 2 && size != 0 { | ||||||
| 		panic("Size must be > 2  (or 0)") | 		panic("Size must be > 2  (or 0)") | ||||||
| 	} | 	} | ||||||
|  |  | ||||||
| 	return &LRUMap{ | 	return &LRUMap[TKey, TData]{ | ||||||
| 		maxsize: size, | 		maxsize: size, | ||||||
| 		lock:    sync.Mutex{}, | 		lock:    sync.Mutex{}, | ||||||
| 		cache:   make(map[string]*cacheNode, size+1), | 		cache:   make(map[TKey]*cacheNode[TKey, TData], size+1), | ||||||
| 		lfuHead: nil, | 		lfuHead: nil, | ||||||
| 		lfuTail: nil, | 		lfuTail: nil, | ||||||
| 	} | 	} | ||||||
| } | } | ||||||
|  |  | ||||||
| func (c *LRUMap) Put(key string, value LRUData) { | func (c *LRUMap[TKey, TData]) Put(key TKey, value TData) { | ||||||
| 	if c.maxsize == 0 { | 	if c.maxsize == 0 { | ||||||
| 		return // cache disabled | 		return // cache disabled | ||||||
| 	} | 	} | ||||||
| @@ -70,7 +68,7 @@ func (c *LRUMap) Put(key string, value LRUData) { | |||||||
| 	} | 	} | ||||||
|  |  | ||||||
| 	// key does not exist: insert into map and add to top of LFU | 	// key does not exist: insert into map and add to top of LFU | ||||||
| 	node = &cacheNode{ | 	node = &cacheNode[TKey, TData]{ | ||||||
| 		key:    key, | 		key:    key, | ||||||
| 		data:   value, | 		data:   value, | ||||||
| 		parent: nil, | 		parent: nil, | ||||||
| @@ -95,9 +93,9 @@ func (c *LRUMap) Put(key string, value LRUData) { | |||||||
| 	} | 	} | ||||||
| } | } | ||||||
|  |  | ||||||
| func (c *LRUMap) TryGet(key string) (LRUData, bool) { | func (c *LRUMap[TKey, TData]) TryGet(key TKey) (TData, bool) { | ||||||
| 	if c.maxsize == 0 { | 	if c.maxsize == 0 { | ||||||
| 		return nil, false // cache disabled | 		return *new(TData), false // cache disabled | ||||||
| 	} | 	} | ||||||
|  |  | ||||||
| 	c.lock.Lock() | 	c.lock.Lock() | ||||||
| @@ -105,13 +103,13 @@ func (c *LRUMap) TryGet(key string) (LRUData, bool) { | |||||||
|  |  | ||||||
| 	val, ok := c.cache[key] | 	val, ok := c.cache[key] | ||||||
| 	if !ok { | 	if !ok { | ||||||
| 		return nil, false | 		return *new(TData), false | ||||||
| 	} | 	} | ||||||
| 	c.moveNodeToTop(val) | 	c.moveNodeToTop(val) | ||||||
| 	return val.data, ok | 	return val.data, ok | ||||||
| } | } | ||||||
|  |  | ||||||
| func (c *LRUMap) moveNodeToTop(node *cacheNode) { | func (c *LRUMap[TKey, TData]) moveNodeToTop(node *cacheNode[TKey, TData]) { | ||||||
| 	// (only called in critical section !) | 	// (only called in critical section !) | ||||||
|  |  | ||||||
| 	if c.lfuHead == node { // fast case | 	if c.lfuHead == node { // fast case | ||||||
| @@ -144,7 +142,7 @@ func (c *LRUMap) moveNodeToTop(node *cacheNode) { | |||||||
| 	} | 	} | ||||||
| } | } | ||||||
|  |  | ||||||
| func (c *LRUMap) Size() int { | func (c *LRUMap[TKey, TData]) Size() int { | ||||||
| 	c.lock.Lock() | 	c.lock.Lock() | ||||||
| 	defer c.lock.Unlock() | 	defer c.lock.Unlock() | ||||||
| 	return len(c.cache) | 	return len(c.cache) | ||||||
|   | |||||||
| @@ -12,7 +12,7 @@ func init() { | |||||||
| } | } | ||||||
|  |  | ||||||
| func TestResultCache1(t *testing.T) { | func TestResultCache1(t *testing.T) { | ||||||
| 	cache := NewLRUMap(8) | 	cache := NewLRUMap[string](8) | ||||||
| 	verifyLRUList(cache, t) | 	verifyLRUList(cache, t) | ||||||
|  |  | ||||||
| 	key := randomKey() | 	key := randomKey() | ||||||
| @@ -39,7 +39,7 @@ func TestResultCache1(t *testing.T) { | |||||||
| 	if !ok { | 	if !ok { | ||||||
| 		t.Errorf("cache TryGet returned no value") | 		t.Errorf("cache TryGet returned no value") | ||||||
| 	} | 	} | ||||||
| 	if !eq(cacheval, val) { | 	if cacheval != val { | ||||||
| 		t.Errorf("cache TryGet returned different value (%+v <> %+v)", cacheval, val) | 		t.Errorf("cache TryGet returned different value (%+v <> %+v)", cacheval, val) | ||||||
| 	} | 	} | ||||||
|  |  | ||||||
| @@ -50,7 +50,7 @@ func TestResultCache1(t *testing.T) { | |||||||
| } | } | ||||||
|  |  | ||||||
| func TestResultCache2(t *testing.T) { | func TestResultCache2(t *testing.T) { | ||||||
| 	cache := NewLRUMap(8) | 	cache := NewLRUMap[string](8) | ||||||
| 	verifyLRUList(cache, t) | 	verifyLRUList(cache, t) | ||||||
|  |  | ||||||
| 	key1 := "key1" | 	key1 := "key1" | ||||||
| @@ -150,7 +150,7 @@ func TestResultCache2(t *testing.T) { | |||||||
| } | } | ||||||
|  |  | ||||||
| func TestResultCache3(t *testing.T) { | func TestResultCache3(t *testing.T) { | ||||||
| 	cache := NewLRUMap(8) | 	cache := NewLRUMap[string](8) | ||||||
| 	verifyLRUList(cache, t) | 	verifyLRUList(cache, t) | ||||||
|  |  | ||||||
| 	key1 := "key1" | 	key1 := "key1" | ||||||
| @@ -160,20 +160,20 @@ func TestResultCache3(t *testing.T) { | |||||||
| 	cache.Put(key1, val1) | 	cache.Put(key1, val1) | ||||||
| 	verifyLRUList(cache, t) | 	verifyLRUList(cache, t) | ||||||
|  |  | ||||||
| 	if val, ok := cache.TryGet(key1); !ok || !eq(val, val1) { | 	if val, ok := cache.TryGet(key1); !ok || val != val1 { | ||||||
| 		t.Errorf("Value in cache should be [val1]") | 		t.Errorf("Value in cache should be [val1]") | ||||||
| 	} | 	} | ||||||
|  |  | ||||||
| 	cache.Put(key1, val2) | 	cache.Put(key1, val2) | ||||||
| 	verifyLRUList(cache, t) | 	verifyLRUList(cache, t) | ||||||
|  |  | ||||||
| 	if val, ok := cache.TryGet(key1); !ok || !eq(val, val2) { | 	if val, ok := cache.TryGet(key1); !ok || val != val2 { | ||||||
| 		t.Errorf("Value in cache should be [val2]") | 		t.Errorf("Value in cache should be [val2]") | ||||||
| 	} | 	} | ||||||
| } | } | ||||||
|  |  | ||||||
| // does a basic consistency check over the internal cache representation | // does a basic consistency check over the internal cache representation | ||||||
| func verifyLRUList(cache *LRUMap, t *testing.T) { | func verifyLRUList[TData any](cache *LRUMap[TData], t *testing.T) { | ||||||
| 	size := 0 | 	size := 0 | ||||||
|  |  | ||||||
| 	tailFound := false | 	tailFound := false | ||||||
| @@ -250,23 +250,10 @@ func randomKey() string { | |||||||
| 	return strconv.FormatInt(rand.Int63(), 16) | 	return strconv.FormatInt(rand.Int63(), 16) | ||||||
| } | } | ||||||
|  |  | ||||||
| func randomVal() LRUData { | func randomVal() string { | ||||||
| 	v, err := langext.NewHexUUID() | 	v, err := langext.NewHexUUID() | ||||||
| 	if err != nil { | 	if err != nil { | ||||||
| 		panic(err) | 		panic(err) | ||||||
| 	} | 	} | ||||||
| 	return &v | 	return v | ||||||
| } |  | ||||||
|  |  | ||||||
| func eq(a LRUData, b LRUData) bool { |  | ||||||
| 	v1, ok1 := a.(*string) |  | ||||||
| 	v2, ok2 := b.(*string) |  | ||||||
| 	if ok1 && ok2 { |  | ||||||
| 		if v1 == nil || v2 == nil { |  | ||||||
| 			return false |  | ||||||
| 		} |  | ||||||
| 		return v1 == v2 |  | ||||||
| 	} |  | ||||||
|  |  | ||||||
| 	return false |  | ||||||
| } | } | ||||||
|   | |||||||
							
								
								
									
										98
									
								
								dataext/stack.go
									
									
									
									
									
										Normal file
									
								
							
							
						
						
									
										98
									
								
								dataext/stack.go
									
									
									
									
									
										Normal file
									
								
							| @@ -0,0 +1,98 @@ | |||||||
|  | package dataext | ||||||
|  |  | ||||||
|  | import ( | ||||||
|  | 	"errors" | ||||||
|  | 	"gogs.mikescher.com/BlackForestBytes/goext/langext" | ||||||
|  | 	"sync" | ||||||
|  | ) | ||||||
|  |  | ||||||
|  | var ErrEmptyStack = errors.New("stack is empty") | ||||||
|  |  | ||||||
|  | type Stack[T any] struct { | ||||||
|  | 	lock *sync.Mutex | ||||||
|  | 	data []T | ||||||
|  | } | ||||||
|  |  | ||||||
|  | func NewStack[T any](threadsafe bool, initialCapacity int) *Stack[T] { | ||||||
|  | 	var lck *sync.Mutex = nil | ||||||
|  | 	if threadsafe { | ||||||
|  | 		lck = &sync.Mutex{} | ||||||
|  | 	} | ||||||
|  | 	return &Stack[T]{ | ||||||
|  | 		lock: lck, | ||||||
|  | 		data: make([]T, 0, initialCapacity), | ||||||
|  | 	} | ||||||
|  | } | ||||||
|  |  | ||||||
|  | func (s *Stack[T]) Push(v T) { | ||||||
|  | 	if s.lock != nil { | ||||||
|  | 		s.lock.Lock() | ||||||
|  | 		defer s.lock.Unlock() | ||||||
|  | 	} | ||||||
|  |  | ||||||
|  | 	s.data = append(s.data, v) | ||||||
|  | } | ||||||
|  |  | ||||||
|  | func (s *Stack[T]) Pop() (T, error) { | ||||||
|  | 	if s.lock != nil { | ||||||
|  | 		s.lock.Lock() | ||||||
|  | 		defer s.lock.Unlock() | ||||||
|  | 	} | ||||||
|  |  | ||||||
|  | 	l := len(s.data) | ||||||
|  | 	if l == 0 { | ||||||
|  | 		return *new(T), ErrEmptyStack | ||||||
|  | 	} | ||||||
|  |  | ||||||
|  | 	result := s.data[l-1] | ||||||
|  | 	s.data = s.data[:l-1] | ||||||
|  |  | ||||||
|  | 	return result, nil | ||||||
|  | } | ||||||
|  |  | ||||||
|  | func (s *Stack[T]) OptPop() *T { | ||||||
|  | 	if s.lock != nil { | ||||||
|  | 		s.lock.Lock() | ||||||
|  | 		defer s.lock.Unlock() | ||||||
|  | 	} | ||||||
|  |  | ||||||
|  | 	l := len(s.data) | ||||||
|  | 	if l == 0 { | ||||||
|  | 		return nil | ||||||
|  | 	} | ||||||
|  |  | ||||||
|  | 	result := s.data[l-1] | ||||||
|  | 	s.data = s.data[:l-1] | ||||||
|  |  | ||||||
|  | 	return langext.Ptr(result) | ||||||
|  | } | ||||||
|  |  | ||||||
|  | func (s *Stack[T]) Peek() (T, error) { | ||||||
|  | 	if s.lock != nil { | ||||||
|  | 		s.lock.Lock() | ||||||
|  | 		defer s.lock.Unlock() | ||||||
|  | 	} | ||||||
|  |  | ||||||
|  | 	l := len(s.data) | ||||||
|  |  | ||||||
|  | 	if l == 0 { | ||||||
|  | 		return *new(T), ErrEmptyStack | ||||||
|  | 	} | ||||||
|  |  | ||||||
|  | 	return s.data[l-1], nil | ||||||
|  | } | ||||||
|  |  | ||||||
|  | func (s *Stack[T]) OptPeek() *T { | ||||||
|  | 	if s.lock != nil { | ||||||
|  | 		s.lock.Lock() | ||||||
|  | 		defer s.lock.Unlock() | ||||||
|  | 	} | ||||||
|  |  | ||||||
|  | 	l := len(s.data) | ||||||
|  |  | ||||||
|  | 	if l == 0 { | ||||||
|  | 		return nil | ||||||
|  | 	} | ||||||
|  |  | ||||||
|  | 	return langext.Ptr(s.data[l-1]) | ||||||
|  | } | ||||||
							
								
								
									
										254
									
								
								dataext/structHash.go
									
									
									
									
									
										Normal file
									
								
							
							
						
						
									
										254
									
								
								dataext/structHash.go
									
									
									
									
									
										Normal file
									
								
							| @@ -0,0 +1,254 @@ | |||||||
|  | package dataext | ||||||
|  |  | ||||||
|  | import ( | ||||||
|  | 	"bytes" | ||||||
|  | 	"crypto/sha256" | ||||||
|  | 	"encoding/binary" | ||||||
|  | 	"errors" | ||||||
|  | 	"fmt" | ||||||
|  | 	"gogs.mikescher.com/BlackForestBytes/goext/langext" | ||||||
|  | 	"hash" | ||||||
|  | 	"io" | ||||||
|  | 	"reflect" | ||||||
|  | 	"sort" | ||||||
|  | ) | ||||||
|  |  | ||||||
|  | type StructHashOptions struct { | ||||||
|  | 	HashAlgo    hash.Hash | ||||||
|  | 	Tag         *string | ||||||
|  | 	SkipChannel bool | ||||||
|  | 	SkipFunc    bool | ||||||
|  | } | ||||||
|  |  | ||||||
|  | func StructHash(dat any, opt ...StructHashOptions) (r []byte, err error) { | ||||||
|  | 	defer func() { | ||||||
|  | 		if rec := recover(); rec != nil { | ||||||
|  | 			r = nil | ||||||
|  | 			err = errors.New(fmt.Sprintf("recovered panic: %v", rec)) | ||||||
|  | 		} | ||||||
|  | 	}() | ||||||
|  |  | ||||||
|  | 	shopt := StructHashOptions{} | ||||||
|  | 	if len(opt) > 1 { | ||||||
|  | 		return nil, errors.New("multiple options supplied") | ||||||
|  | 	} else if len(opt) == 1 { | ||||||
|  | 		shopt = opt[0] | ||||||
|  | 	} | ||||||
|  |  | ||||||
|  | 	if shopt.HashAlgo == nil { | ||||||
|  | 		shopt.HashAlgo = sha256.New() | ||||||
|  | 	} | ||||||
|  |  | ||||||
|  | 	writer := new(bytes.Buffer) | ||||||
|  |  | ||||||
|  | 	if langext.IsNil(dat) { | ||||||
|  | 		shopt.HashAlgo.Reset() | ||||||
|  | 		shopt.HashAlgo.Write(writer.Bytes()) | ||||||
|  | 		res := shopt.HashAlgo.Sum(nil) | ||||||
|  | 		return res, nil | ||||||
|  | 	} | ||||||
|  |  | ||||||
|  | 	err = binarize(writer, reflect.ValueOf(dat), shopt) | ||||||
|  | 	if err != nil { | ||||||
|  | 		return nil, err | ||||||
|  | 	} | ||||||
|  |  | ||||||
|  | 	shopt.HashAlgo.Reset() | ||||||
|  | 	shopt.HashAlgo.Write(writer.Bytes()) | ||||||
|  | 	res := shopt.HashAlgo.Sum(nil) | ||||||
|  |  | ||||||
|  | 	return res, nil | ||||||
|  | } | ||||||
|  |  | ||||||
|  | func writeBinarized(writer io.Writer, dat any) error { | ||||||
|  | 	tmp := bytes.Buffer{} | ||||||
|  | 	err := binary.Write(&tmp, binary.LittleEndian, dat) | ||||||
|  | 	if err != nil { | ||||||
|  | 		return err | ||||||
|  | 	} | ||||||
|  | 	err = binary.Write(writer, binary.LittleEndian, uint64(tmp.Len())) | ||||||
|  | 	if err != nil { | ||||||
|  | 		return err | ||||||
|  | 	} | ||||||
|  | 	_, err = writer.Write(tmp.Bytes()) | ||||||
|  | 	if err != nil { | ||||||
|  | 		return err | ||||||
|  | 	} | ||||||
|  | 	return nil | ||||||
|  | } | ||||||
|  |  | ||||||
|  | func binarize(writer io.Writer, dat reflect.Value, opt StructHashOptions) error { | ||||||
|  | 	var err error | ||||||
|  |  | ||||||
|  | 	err = binary.Write(writer, binary.LittleEndian, uint8(dat.Kind())) | ||||||
|  | 	switch dat.Kind() { | ||||||
|  | 	case reflect.Ptr, reflect.Map, reflect.Array, reflect.Chan, reflect.Slice, reflect.Interface: | ||||||
|  | 		if dat.IsNil() { | ||||||
|  | 			err = binary.Write(writer, binary.LittleEndian, uint64(0)) | ||||||
|  | 			if err != nil { | ||||||
|  | 				return err | ||||||
|  | 			} | ||||||
|  | 			return nil | ||||||
|  | 		} | ||||||
|  | 	} | ||||||
|  |  | ||||||
|  | 	err = binary.Write(writer, binary.LittleEndian, uint64(len(dat.Type().String()))) | ||||||
|  | 	if err != nil { | ||||||
|  | 		return err | ||||||
|  | 	} | ||||||
|  | 	_, err = writer.Write([]byte(dat.Type().String())) | ||||||
|  | 	if err != nil { | ||||||
|  | 		return err | ||||||
|  | 	} | ||||||
|  |  | ||||||
|  | 	switch dat.Type().Kind() { | ||||||
|  | 	case reflect.Invalid: | ||||||
|  | 		return errors.New("cannot binarize value of kind <Invalid>") | ||||||
|  | 	case reflect.Bool: | ||||||
|  | 		return writeBinarized(writer, dat.Bool()) | ||||||
|  | 	case reflect.Int: | ||||||
|  | 		return writeBinarized(writer, int64(dat.Int())) | ||||||
|  | 	case reflect.Int8: | ||||||
|  | 		fallthrough | ||||||
|  | 	case reflect.Int16: | ||||||
|  | 		fallthrough | ||||||
|  | 	case reflect.Int32: | ||||||
|  | 		fallthrough | ||||||
|  | 	case reflect.Int64: | ||||||
|  | 		return writeBinarized(writer, dat.Interface()) | ||||||
|  | 	case reflect.Uint: | ||||||
|  | 		return writeBinarized(writer, uint64(dat.Int())) | ||||||
|  | 	case reflect.Uint8: | ||||||
|  | 		fallthrough | ||||||
|  | 	case reflect.Uint16: | ||||||
|  | 		fallthrough | ||||||
|  | 	case reflect.Uint32: | ||||||
|  | 		fallthrough | ||||||
|  | 	case reflect.Uint64: | ||||||
|  | 		return writeBinarized(writer, dat.Interface()) | ||||||
|  | 	case reflect.Uintptr: | ||||||
|  | 		return errors.New("cannot binarize value of kind <Uintptr>") | ||||||
|  | 	case reflect.Float32: | ||||||
|  | 		fallthrough | ||||||
|  | 	case reflect.Float64: | ||||||
|  | 		return writeBinarized(writer, dat.Interface()) | ||||||
|  | 	case reflect.Complex64: | ||||||
|  | 		return errors.New("cannot binarize value of kind <Complex64>") | ||||||
|  | 	case reflect.Complex128: | ||||||
|  | 		return errors.New("cannot binarize value of kind <Complex128>") | ||||||
|  | 	case reflect.Slice: | ||||||
|  | 		fallthrough | ||||||
|  | 	case reflect.Array: | ||||||
|  | 		return binarizeArrayOrSlice(writer, dat, opt) | ||||||
|  | 	case reflect.Chan: | ||||||
|  | 		if opt.SkipChannel { | ||||||
|  | 			return nil | ||||||
|  | 		} | ||||||
|  | 		return errors.New("cannot binarize value of kind <Chan>") | ||||||
|  | 	case reflect.Func: | ||||||
|  | 		if opt.SkipFunc { | ||||||
|  | 			return nil | ||||||
|  | 		} | ||||||
|  | 		return errors.New("cannot binarize value of kind <Func>") | ||||||
|  | 	case reflect.Interface: | ||||||
|  | 		return binarize(writer, dat.Elem(), opt) | ||||||
|  | 	case reflect.Map: | ||||||
|  | 		return binarizeMap(writer, dat, opt) | ||||||
|  | 	case reflect.Pointer: | ||||||
|  | 		return binarize(writer, dat.Elem(), opt) | ||||||
|  | 	case reflect.String: | ||||||
|  | 		v := dat.String() | ||||||
|  | 		err = binary.Write(writer, binary.LittleEndian, uint64(len(v))) | ||||||
|  | 		if err != nil { | ||||||
|  | 			return err | ||||||
|  | 		} | ||||||
|  | 		_, err = writer.Write([]byte(v)) | ||||||
|  | 		if err != nil { | ||||||
|  | 			return err | ||||||
|  | 		} | ||||||
|  | 		return nil | ||||||
|  | 	case reflect.Struct: | ||||||
|  | 		return binarizeStruct(writer, dat, opt) | ||||||
|  | 	case reflect.UnsafePointer: | ||||||
|  | 		return errors.New("cannot binarize value of kind <UnsafePointer>") | ||||||
|  | 	default: | ||||||
|  | 		return errors.New("cannot binarize value of unknown kind <" + dat.Type().Kind().String() + ">") | ||||||
|  | 	} | ||||||
|  | } | ||||||
|  |  | ||||||
|  | func binarizeStruct(writer io.Writer, dat reflect.Value, opt StructHashOptions) error { | ||||||
|  | 	err := binary.Write(writer, binary.LittleEndian, uint64(dat.NumField())) | ||||||
|  | 	if err != nil { | ||||||
|  | 		return err | ||||||
|  | 	} | ||||||
|  |  | ||||||
|  | 	for i := 0; i < dat.NumField(); i++ { | ||||||
|  |  | ||||||
|  | 		if opt.Tag != nil { | ||||||
|  | 			if _, ok := dat.Type().Field(i).Tag.Lookup(*opt.Tag); !ok { | ||||||
|  | 				continue | ||||||
|  | 			} | ||||||
|  | 		} | ||||||
|  |  | ||||||
|  | 		err = binary.Write(writer, binary.LittleEndian, uint64(len(dat.Type().Field(i).Name))) | ||||||
|  | 		if err != nil { | ||||||
|  | 			return err | ||||||
|  | 		} | ||||||
|  | 		_, err = writer.Write([]byte(dat.Type().Field(i).Name)) | ||||||
|  | 		if err != nil { | ||||||
|  | 			return err | ||||||
|  | 		} | ||||||
|  |  | ||||||
|  | 		err = binarize(writer, dat.Field(i), opt) | ||||||
|  | 		if err != nil { | ||||||
|  | 			return err | ||||||
|  | 		} | ||||||
|  | 	} | ||||||
|  |  | ||||||
|  | 	return nil | ||||||
|  | } | ||||||
|  |  | ||||||
|  | func binarizeArrayOrSlice(writer io.Writer, dat reflect.Value, opt StructHashOptions) error { | ||||||
|  | 	err := binary.Write(writer, binary.LittleEndian, uint64(dat.Len())) | ||||||
|  | 	if err != nil { | ||||||
|  | 		return err | ||||||
|  | 	} | ||||||
|  |  | ||||||
|  | 	for i := 0; i < dat.Len(); i++ { | ||||||
|  | 		err := binarize(writer, dat.Index(i), opt) | ||||||
|  | 		if err != nil { | ||||||
|  | 			return err | ||||||
|  | 		} | ||||||
|  | 	} | ||||||
|  |  | ||||||
|  | 	return nil | ||||||
|  | } | ||||||
|  |  | ||||||
|  | func binarizeMap(writer io.Writer, dat reflect.Value, opt StructHashOptions) error { | ||||||
|  | 	err := binary.Write(writer, binary.LittleEndian, uint64(dat.Len())) | ||||||
|  | 	if err != nil { | ||||||
|  | 		return err | ||||||
|  | 	} | ||||||
|  |  | ||||||
|  | 	sub := make([][]byte, 0, dat.Len()) | ||||||
|  |  | ||||||
|  | 	for _, k := range dat.MapKeys() { | ||||||
|  | 		tmp := bytes.Buffer{} | ||||||
|  | 		err = binarize(&tmp, dat.MapIndex(k), opt) | ||||||
|  | 		if err != nil { | ||||||
|  | 			return err | ||||||
|  | 		} | ||||||
|  | 		sub = append(sub, tmp.Bytes()) | ||||||
|  | 	} | ||||||
|  |  | ||||||
|  | 	sort.Slice(sub, func(i1, i2 int) bool { return bytes.Compare(sub[i1], sub[i2]) < 0 }) | ||||||
|  |  | ||||||
|  | 	for _, v := range sub { | ||||||
|  | 		_, err = writer.Write(v) | ||||||
|  | 		if err != nil { | ||||||
|  | 			return err | ||||||
|  | 		} | ||||||
|  | 	} | ||||||
|  |  | ||||||
|  | 	return nil | ||||||
|  | } | ||||||
							
								
								
									
										143
									
								
								dataext/structHash_test.go
									
									
									
									
									
										Normal file
									
								
							
							
						
						
									
										143
									
								
								dataext/structHash_test.go
									
									
									
									
									
										Normal file
									
								
							| @@ -0,0 +1,143 @@ | |||||||
|  | package dataext | ||||||
|  |  | ||||||
|  | import ( | ||||||
|  | 	"encoding/hex" | ||||||
|  | 	"gogs.mikescher.com/BlackForestBytes/goext/langext" | ||||||
|  | 	"testing" | ||||||
|  | ) | ||||||
|  |  | ||||||
|  | func noErrStructHash(t *testing.T, dat any, opt ...StructHashOptions) []byte { | ||||||
|  | 	res, err := StructHash(dat, opt...) | ||||||
|  | 	if err != nil { | ||||||
|  | 		t.Error(err) | ||||||
|  | 		t.FailNow() | ||||||
|  | 		return nil | ||||||
|  | 	} | ||||||
|  | 	return res | ||||||
|  | } | ||||||
|  |  | ||||||
|  | func TestStructHashSimple(t *testing.T) { | ||||||
|  |  | ||||||
|  | 	assertEqual(t, "209bf774af36cc3a045c152d9f1269ef3684ad819c1359ee73ff0283a308fefa", noErrStructHash(t, "Hello")) | ||||||
|  | 	assertEqual(t, "c32f3626b981ae2997db656f3acad3f1dc9d30ef6b6d14296c023e391b25f71a", noErrStructHash(t, 0)) | ||||||
|  | 	assertEqual(t, "01b781b03e9586b257d387057dfc70d9f06051e7d3c1e709a57e13cc8daf3e35", noErrStructHash(t, []byte{})) | ||||||
|  | 	assertEqual(t, "93e1dcd45c732fe0079b0fb3204c7c803f0921835f6bfee2e6ff263e73eed53c", noErrStructHash(t, []int{})) | ||||||
|  | 	assertEqual(t, "54f637a376aad55b3160d98ebbcae8099b70d91b9400df23fb3709855d59800a", noErrStructHash(t, []int{1, 2, 3})) | ||||||
|  | 	assertEqual(t, "e3b0c44298fc1c149afbf4c8996fb92427ae41e4649b934ca495991b7852b855", noErrStructHash(t, nil)) | ||||||
|  | 	assertEqual(t, "349a7db91aa78fd30bbaa7c7f9c7bfb2fcfe72869b4861162a96713a852f60d3", noErrStructHash(t, []any{1, "", nil})) | ||||||
|  | 	assertEqual(t, "ca51aab87808bf0062a4a024de6aac0c2bad54275cc857a4944569f89fd245ad", noErrStructHash(t, struct{}{})) | ||||||
|  |  | ||||||
|  | } | ||||||
|  |  | ||||||
|  | func TestStructHashSimpleStruct(t *testing.T) { | ||||||
|  |  | ||||||
|  | 	type t0 struct { | ||||||
|  | 		F1 int | ||||||
|  | 		F2 []string | ||||||
|  | 		F3 *int | ||||||
|  | 	} | ||||||
|  |  | ||||||
|  | 	assertEqual(t, "a90bff751c70c738bb5cfc9b108e783fa9c19c0bc9273458e0aaee6e74aa1b92", noErrStructHash(t, t0{ | ||||||
|  | 		F1: 10, | ||||||
|  | 		F2: []string{"1", "2", "3"}, | ||||||
|  | 		F3: nil, | ||||||
|  | 	})) | ||||||
|  |  | ||||||
|  | 	assertEqual(t, "5d09090dc34ac59dd645f197a255f653387723de3afa1b614721ea5a081c675f", noErrStructHash(t, t0{ | ||||||
|  | 		F1: 10, | ||||||
|  | 		F2: []string{"1", "2", "3"}, | ||||||
|  | 		F3: langext.Ptr(99), | ||||||
|  | 	})) | ||||||
|  |  | ||||||
|  | } | ||||||
|  |  | ||||||
|  | func TestStructHashLayeredStruct(t *testing.T) { | ||||||
|  |  | ||||||
|  | 	type t1_1 struct { | ||||||
|  | 		F10 float32 | ||||||
|  | 		F12 float64 | ||||||
|  | 		F15 bool | ||||||
|  | 	} | ||||||
|  | 	type t1_2 struct { | ||||||
|  | 		SV1 *t1_1 | ||||||
|  | 		SV2 *t1_1 | ||||||
|  | 		SV3 t1_1 | ||||||
|  | 	} | ||||||
|  |  | ||||||
|  | 	assertEqual(t, "fd4ca071fb40a288fee4b7a3dfdaab577b30cb8f80f81ec511e7afd72dc3b469", noErrStructHash(t, t1_2{ | ||||||
|  | 		SV1: nil, | ||||||
|  | 		SV2: nil, | ||||||
|  | 		SV3: t1_1{ | ||||||
|  | 			F10: 1, | ||||||
|  | 			F12: 2, | ||||||
|  | 			F15: false, | ||||||
|  | 		}, | ||||||
|  | 	})) | ||||||
|  | 	assertEqual(t, "3fbf7c67d8121deda075cc86319a4e32d71744feb2cebf89b43bc682f072a029", noErrStructHash(t, t1_2{ | ||||||
|  | 		SV1: nil, | ||||||
|  | 		SV2: &t1_1{}, | ||||||
|  | 		SV3: t1_1{ | ||||||
|  | 			F10: 3, | ||||||
|  | 			F12: 4, | ||||||
|  | 			F15: true, | ||||||
|  | 		}, | ||||||
|  | 	})) | ||||||
|  | 	assertEqual(t, "b1791ccd1b346c3ede5bbffda85555adcd8216b93ffca23f14fe175ec47c5104", noErrStructHash(t, t1_2{ | ||||||
|  | 		SV1: &t1_1{}, | ||||||
|  | 		SV2: &t1_1{}, | ||||||
|  | 		SV3: t1_1{ | ||||||
|  | 			F10: 5, | ||||||
|  | 			F12: 6, | ||||||
|  | 			F15: false, | ||||||
|  | 		}, | ||||||
|  | 	})) | ||||||
|  |  | ||||||
|  | } | ||||||
|  |  | ||||||
|  | func TestStructHashMap(t *testing.T) { | ||||||
|  |  | ||||||
|  | 	type t0 struct { | ||||||
|  | 		F1 int | ||||||
|  | 		F2 map[string]int | ||||||
|  | 	} | ||||||
|  |  | ||||||
|  | 	assertEqual(t, "d50c53ad1fafb448c33fddd5aca01a86a2edf669ce2ecab07ba6fe877951d824", noErrStructHash(t, t0{ | ||||||
|  | 		F1: 10, | ||||||
|  | 		F2: map[string]int{ | ||||||
|  | 			"x": 1, | ||||||
|  | 			"0": 2, | ||||||
|  | 			"a": 99, | ||||||
|  | 		}, | ||||||
|  | 	})) | ||||||
|  |  | ||||||
|  | 	assertEqual(t, "d50c53ad1fafb448c33fddd5aca01a86a2edf669ce2ecab07ba6fe877951d824", noErrStructHash(t, t0{ | ||||||
|  | 		F1: 10, | ||||||
|  | 		F2: map[string]int{ | ||||||
|  | 			"a": 99, | ||||||
|  | 			"x": 1, | ||||||
|  | 			"0": 2, | ||||||
|  | 		}, | ||||||
|  | 	})) | ||||||
|  |  | ||||||
|  | 	m3 := make(map[string]int, 99) | ||||||
|  | 	m3["a"] = 0 | ||||||
|  | 	m3["x"] = 0 | ||||||
|  | 	m3["0"] = 0 | ||||||
|  |  | ||||||
|  | 	m3["0"] = 99 | ||||||
|  | 	m3["x"] = 1 | ||||||
|  | 	m3["a"] = 2 | ||||||
|  |  | ||||||
|  | 	assertEqual(t, "d50c53ad1fafb448c33fddd5aca01a86a2edf669ce2ecab07ba6fe877951d824", noErrStructHash(t, t0{ | ||||||
|  | 		F1: 10, | ||||||
|  | 		F2: m3, | ||||||
|  | 	})) | ||||||
|  |  | ||||||
|  | } | ||||||
|  |  | ||||||
|  | func assertEqual(t *testing.T, expected string, actual []byte) { | ||||||
|  | 	actualStr := hex.EncodeToString(actual) | ||||||
|  | 	if actualStr != expected { | ||||||
|  | 		t.Errorf("values differ: Actual: '%v', Expected: '%v'", actualStr, expected) | ||||||
|  | 	} | ||||||
|  | } | ||||||
| @@ -2,17 +2,17 @@ package dataext | |||||||
|  |  | ||||||
| import "sync" | import "sync" | ||||||
|  |  | ||||||
| type SyncStringSet struct { | type SyncSet[TData comparable] struct { | ||||||
| 	data map[string]bool | 	data map[TData]bool | ||||||
| 	lock sync.Mutex | 	lock sync.Mutex | ||||||
| } | } | ||||||
|  |  | ||||||
| func (s *SyncStringSet) Add(value string) bool { | func (s *SyncSet[TData]) Add(value TData) bool { | ||||||
| 	s.lock.Lock() | 	s.lock.Lock() | ||||||
| 	defer s.lock.Unlock() | 	defer s.lock.Unlock() | ||||||
|  |  | ||||||
| 	if s.data == nil { | 	if s.data == nil { | ||||||
| 		s.data = make(map[string]bool) | 		s.data = make(map[TData]bool) | ||||||
| 	} | 	} | ||||||
|  |  | ||||||
| 	_, ok := s.data[value] | 	_, ok := s.data[value] | ||||||
| @@ -21,12 +21,12 @@ func (s *SyncStringSet) Add(value string) bool { | |||||||
| 	return !ok | 	return !ok | ||||||
| } | } | ||||||
|  |  | ||||||
| func (s *SyncStringSet) AddAll(values []string) { | func (s *SyncSet[TData]) AddAll(values []TData) { | ||||||
| 	s.lock.Lock() | 	s.lock.Lock() | ||||||
| 	defer s.lock.Unlock() | 	defer s.lock.Unlock() | ||||||
|  |  | ||||||
| 	if s.data == nil { | 	if s.data == nil { | ||||||
| 		s.data = make(map[string]bool) | 		s.data = make(map[TData]bool) | ||||||
| 	} | 	} | ||||||
|  |  | ||||||
| 	for _, value := range values { | 	for _, value := range values { | ||||||
| @@ -34,12 +34,12 @@ func (s *SyncStringSet) AddAll(values []string) { | |||||||
| 	} | 	} | ||||||
| } | } | ||||||
|  |  | ||||||
| func (s *SyncStringSet) Contains(value string) bool { | func (s *SyncSet[TData]) Contains(value TData) bool { | ||||||
| 	s.lock.Lock() | 	s.lock.Lock() | ||||||
| 	defer s.lock.Unlock() | 	defer s.lock.Unlock() | ||||||
|  |  | ||||||
| 	if s.data == nil { | 	if s.data == nil { | ||||||
| 		s.data = make(map[string]bool) | 		s.data = make(map[TData]bool) | ||||||
| 	} | 	} | ||||||
|  |  | ||||||
| 	_, ok := s.data[value] | 	_, ok := s.data[value] | ||||||
| @@ -47,15 +47,15 @@ func (s *SyncStringSet) Contains(value string) bool { | |||||||
| 	return ok | 	return ok | ||||||
| } | } | ||||||
|  |  | ||||||
| func (s *SyncStringSet) Get() []string { | func (s *SyncSet[TData]) Get() []TData { | ||||||
| 	s.lock.Lock() | 	s.lock.Lock() | ||||||
| 	defer s.lock.Unlock() | 	defer s.lock.Unlock() | ||||||
|  |  | ||||||
| 	if s.data == nil { | 	if s.data == nil { | ||||||
| 		s.data = make(map[string]bool) | 		s.data = make(map[TData]bool) | ||||||
| 	} | 	} | ||||||
|  |  | ||||||
| 	r := make([]string, 0, len(s.data)) | 	r := make([]TData, 0, len(s.data)) | ||||||
|  |  | ||||||
| 	for k := range s.data { | 	for k := range s.data { | ||||||
| 		r = append(r, k) | 		r = append(r, k) | ||||||
|   | |||||||
							
								
								
									
										9
									
								
								go.mod
									
									
									
									
									
								
							
							
						
						
									
										9
									
								
								go.mod
									
									
									
									
									
								
							| @@ -3,6 +3,11 @@ module gogs.mikescher.com/BlackForestBytes/goext | |||||||
| go 1.19 | go 1.19 | ||||||
|  |  | ||||||
| require ( | require ( | ||||||
| 	golang.org/x/sys v0.1.0 | 	golang.org/x/sys v0.3.0 | ||||||
| 	golang.org/x/term v0.1.0 | 	golang.org/x/term v0.3.0 | ||||||
|  | ) | ||||||
|  |  | ||||||
|  | require ( | ||||||
|  | 	github.com/jmoiron/sqlx v1.3.5 // indirect | ||||||
|  | 	golang.org/x/crypto v0.4.0 // indirect | ||||||
| ) | ) | ||||||
|   | |||||||
							
								
								
									
										11
									
								
								go.sum
									
									
									
									
									
								
							
							
						
						
									
										11
									
								
								go.sum
									
									
									
									
									
								
							| @@ -1,4 +1,15 @@ | |||||||
|  | github.com/go-sql-driver/mysql v1.6.0/go.mod h1:DCzpHaOWr8IXmIStZouvnhqoel9Qv2LBy8hT2VhHyBg= | ||||||
|  | github.com/jmoiron/sqlx v1.3.5 h1:vFFPA71p1o5gAeqtEAwLU4dnX2napprKtHr7PYIcN3g= | ||||||
|  | github.com/jmoiron/sqlx v1.3.5/go.mod h1:nRVWtLre0KfCLJvgxzCsLVMogSvQ1zNJtpYr2Ccp0mQ= | ||||||
|  | github.com/lib/pq v1.2.0/go.mod h1:5WUZQaWbwv1U+lTReE5YruASi9Al49XbQIvNi/34Woo= | ||||||
|  | github.com/mattn/go-sqlite3 v1.14.6/go.mod h1:NyWgC/yNuGj7Q9rpYnZvas74GogHl5/Z4A/KQRfk6bU= | ||||||
|  | golang.org/x/crypto v0.4.0 h1:UVQgzMY87xqpKNgb+kDsll2Igd33HszWHFLmpaRMq/8= | ||||||
|  | golang.org/x/crypto v0.4.0/go.mod h1:3quD/ATkf6oY+rnes5c3ExXTbLc8mueNue5/DoinL80= | ||||||
| golang.org/x/sys v0.1.0 h1:kunALQeHf1/185U1i0GOB/fy1IPRDDpuoOOqRReG57U= | golang.org/x/sys v0.1.0 h1:kunALQeHf1/185U1i0GOB/fy1IPRDDpuoOOqRReG57U= | ||||||
| golang.org/x/sys v0.1.0/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg= | golang.org/x/sys v0.1.0/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg= | ||||||
|  | golang.org/x/sys v0.3.0 h1:w8ZOecv6NaNa/zC8944JTU3vz4u6Lagfk4RPQxv92NQ= | ||||||
|  | golang.org/x/sys v0.3.0/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg= | ||||||
| golang.org/x/term v0.1.0 h1:g6Z6vPFA9dYBAF7DWcH6sCcOntplXsDKcliusYijMlw= | golang.org/x/term v0.1.0 h1:g6Z6vPFA9dYBAF7DWcH6sCcOntplXsDKcliusYijMlw= | ||||||
| golang.org/x/term v0.1.0/go.mod h1:jbD1KX2456YbFQfuXm/mYQcufACuNUgVhRMnK/tPxf8= | golang.org/x/term v0.1.0/go.mod h1:jbD1KX2456YbFQfuXm/mYQcufACuNUgVhRMnK/tPxf8= | ||||||
|  | golang.org/x/term v0.3.0 h1:qoo4akIqOcDME5bhc/NgxUdovd6BSS2uMsVjB56q1xI= | ||||||
|  | golang.org/x/term v0.3.0/go.mod h1:q750SLmJuPmVoN1blW3UFBPREJfb1KmY3vwxfr+nFDA= | ||||||
|   | |||||||
| @@ -70,7 +70,73 @@ func ArrEqualsExact[T comparable](arr1 []T, arr2 []T) bool { | |||||||
| 	return true | 	return true | ||||||
| } | } | ||||||
|  |  | ||||||
| func ArrAll(arr interface{}, fn func(int) bool) bool { | func ArrAll[T any](arr []T, fn func(T) bool) bool { | ||||||
|  | 	for _, av := range arr { | ||||||
|  | 		if !fn(av) { | ||||||
|  | 			return false | ||||||
|  | 		} | ||||||
|  | 	} | ||||||
|  | 	return true | ||||||
|  | } | ||||||
|  |  | ||||||
|  | func ArrAllErr[T any](arr []T, fn func(T) (bool, error)) (bool, error) { | ||||||
|  | 	for _, av := range arr { | ||||||
|  | 		v, err := fn(av) | ||||||
|  | 		if err != nil { | ||||||
|  | 			return false, err | ||||||
|  | 		} | ||||||
|  | 		if !v { | ||||||
|  | 			return false, nil | ||||||
|  | 		} | ||||||
|  | 	} | ||||||
|  | 	return true, nil | ||||||
|  | } | ||||||
|  |  | ||||||
|  | func ArrNone[T any](arr []T, fn func(T) bool) bool { | ||||||
|  | 	for _, av := range arr { | ||||||
|  | 		if fn(av) { | ||||||
|  | 			return false | ||||||
|  | 		} | ||||||
|  | 	} | ||||||
|  | 	return true | ||||||
|  | } | ||||||
|  |  | ||||||
|  | func ArrNoneErr[T any](arr []T, fn func(T) (bool, error)) (bool, error) { | ||||||
|  | 	for _, av := range arr { | ||||||
|  | 		v, err := fn(av) | ||||||
|  | 		if err != nil { | ||||||
|  | 			return false, err | ||||||
|  | 		} | ||||||
|  | 		if v { | ||||||
|  | 			return false, nil | ||||||
|  | 		} | ||||||
|  | 	} | ||||||
|  | 	return true, nil | ||||||
|  | } | ||||||
|  |  | ||||||
|  | func ArrAny[T any](arr []T, fn func(T) bool) bool { | ||||||
|  | 	for _, av := range arr { | ||||||
|  | 		if fn(av) { | ||||||
|  | 			return true | ||||||
|  | 		} | ||||||
|  | 	} | ||||||
|  | 	return false | ||||||
|  | } | ||||||
|  |  | ||||||
|  | func ArrAnyErr[T any](arr []T, fn func(T) (bool, error)) (bool, error) { | ||||||
|  | 	for _, av := range arr { | ||||||
|  | 		v, err := fn(av) | ||||||
|  | 		if err != nil { | ||||||
|  | 			return false, err | ||||||
|  | 		} | ||||||
|  | 		if v { | ||||||
|  | 			return true, nil | ||||||
|  | 		} | ||||||
|  | 	} | ||||||
|  | 	return false, nil | ||||||
|  | } | ||||||
|  |  | ||||||
|  | func ArrIdxAll(arr any, fn func(int) bool) bool { | ||||||
| 	av := reflect.ValueOf(arr) | 	av := reflect.ValueOf(arr) | ||||||
| 	for i := 0; i < av.Len(); i++ { | 	for i := 0; i < av.Len(); i++ { | ||||||
| 		if !fn(i) { | 		if !fn(i) { | ||||||
| @@ -80,7 +146,7 @@ func ArrAll(arr interface{}, fn func(int) bool) bool { | |||||||
| 	return true | 	return true | ||||||
| } | } | ||||||
|  |  | ||||||
| func ArrAllErr(arr interface{}, fn func(int) (bool, error)) (bool, error) { | func ArrIdxAllErr(arr any, fn func(int) (bool, error)) (bool, error) { | ||||||
| 	av := reflect.ValueOf(arr) | 	av := reflect.ValueOf(arr) | ||||||
| 	for i := 0; i < av.Len(); i++ { | 	for i := 0; i < av.Len(); i++ { | ||||||
| 		v, err := fn(i) | 		v, err := fn(i) | ||||||
| @@ -94,7 +160,7 @@ func ArrAllErr(arr interface{}, fn func(int) (bool, error)) (bool, error) { | |||||||
| 	return true, nil | 	return true, nil | ||||||
| } | } | ||||||
|  |  | ||||||
| func ArrNone(arr interface{}, fn func(int) bool) bool { | func ArrIdxNone(arr any, fn func(int) bool) bool { | ||||||
| 	av := reflect.ValueOf(arr) | 	av := reflect.ValueOf(arr) | ||||||
| 	for i := 0; i < av.Len(); i++ { | 	for i := 0; i < av.Len(); i++ { | ||||||
| 		if fn(i) { | 		if fn(i) { | ||||||
| @@ -104,7 +170,7 @@ func ArrNone(arr interface{}, fn func(int) bool) bool { | |||||||
| 	return true | 	return true | ||||||
| } | } | ||||||
|  |  | ||||||
| func ArrNoneErr(arr interface{}, fn func(int) (bool, error)) (bool, error) { | func ArrIdxNoneErr(arr any, fn func(int) (bool, error)) (bool, error) { | ||||||
| 	av := reflect.ValueOf(arr) | 	av := reflect.ValueOf(arr) | ||||||
| 	for i := 0; i < av.Len(); i++ { | 	for i := 0; i < av.Len(); i++ { | ||||||
| 		v, err := fn(i) | 		v, err := fn(i) | ||||||
| @@ -118,7 +184,7 @@ func ArrNoneErr(arr interface{}, fn func(int) (bool, error)) (bool, error) { | |||||||
| 	return true, nil | 	return true, nil | ||||||
| } | } | ||||||
|  |  | ||||||
| func ArrAny(arr interface{}, fn func(int) bool) bool { | func ArrIdxAny(arr any, fn func(int) bool) bool { | ||||||
| 	av := reflect.ValueOf(arr) | 	av := reflect.ValueOf(arr) | ||||||
| 	for i := 0; i < av.Len(); i++ { | 	for i := 0; i < av.Len(); i++ { | ||||||
| 		if fn(i) { | 		if fn(i) { | ||||||
| @@ -128,7 +194,7 @@ func ArrAny(arr interface{}, fn func(int) bool) bool { | |||||||
| 	return false | 	return false | ||||||
| } | } | ||||||
|  |  | ||||||
| func ArrAnyErr(arr interface{}, fn func(int) (bool, error)) (bool, error) { | func ArrIdxAnyErr(arr any, fn func(int) (bool, error)) (bool, error) { | ||||||
| 	av := reflect.ValueOf(arr) | 	av := reflect.ValueOf(arr) | ||||||
| 	for i := 0; i < av.Len(); i++ { | 	for i := 0; i < av.Len(); i++ { | ||||||
| 		v, err := fn(i) | 		v, err := fn(i) | ||||||
| @@ -142,7 +208,7 @@ func ArrAnyErr(arr interface{}, fn func(int) (bool, error)) (bool, error) { | |||||||
| 	return false, nil | 	return false, nil | ||||||
| } | } | ||||||
|  |  | ||||||
| func ArrFirst[T comparable](arr []T, comp func(v T) bool) (T, bool) { | func ArrFirst[T any](arr []T, comp func(v T) bool) (T, bool) { | ||||||
| 	for _, v := range arr { | 	for _, v := range arr { | ||||||
| 		if comp(v) { | 		if comp(v) { | ||||||
| 			return v, true | 			return v, true | ||||||
| @@ -151,7 +217,7 @@ func ArrFirst[T comparable](arr []T, comp func(v T) bool) (T, bool) { | |||||||
| 	return *new(T), false | 	return *new(T), false | ||||||
| } | } | ||||||
|  |  | ||||||
| func ArrLast[T comparable](arr []T, comp func(v T) bool) (T, bool) { | func ArrLast[T any](arr []T, comp func(v T) bool) (T, bool) { | ||||||
| 	found := false | 	found := false | ||||||
| 	result := *new(T) | 	result := *new(T) | ||||||
| 	for _, v := range arr { | 	for _, v := range arr { | ||||||
|   | |||||||
| @@ -15,3 +15,35 @@ func Conditional[T any](v bool, resTrue T, resFalse T) T { | |||||||
| 		return resFalse | 		return resFalse | ||||||
| 	} | 	} | ||||||
| } | } | ||||||
|  |  | ||||||
|  | func ConditionalFn00[T any](v bool, resTrue T, resFalse T) T { | ||||||
|  | 	if v { | ||||||
|  | 		return resTrue | ||||||
|  | 	} else { | ||||||
|  | 		return resFalse | ||||||
|  | 	} | ||||||
|  | } | ||||||
|  |  | ||||||
|  | func ConditionalFn10[T any](v bool, resTrue func() T, resFalse T) T { | ||||||
|  | 	if v { | ||||||
|  | 		return resTrue() | ||||||
|  | 	} else { | ||||||
|  | 		return resFalse | ||||||
|  | 	} | ||||||
|  | } | ||||||
|  |  | ||||||
|  | func ConditionalFn01[T any](v bool, resTrue T, resFalse func() T) T { | ||||||
|  | 	if v { | ||||||
|  | 		return resTrue | ||||||
|  | 	} else { | ||||||
|  | 		return resFalse() | ||||||
|  | 	} | ||||||
|  | } | ||||||
|  |  | ||||||
|  | func ConditionalFn11[T any](v bool, resTrue func() T, resFalse func() T) T { | ||||||
|  | 	if v { | ||||||
|  | 		return resTrue() | ||||||
|  | 	} else { | ||||||
|  | 		return resFalse() | ||||||
|  | 	} | ||||||
|  | } | ||||||
|   | |||||||
| @@ -34,3 +34,7 @@ func IsNil(i interface{}) bool { | |||||||
| 	} | 	} | ||||||
| 	return false | 	return false | ||||||
| } | } | ||||||
|  |  | ||||||
|  | func PtrEquals[T comparable](v1 *T, v2 *T) bool { | ||||||
|  | 	return (v1 == nil && v2 == nil) || (v1 != nil && v2 != nil && *v1 == *v2) | ||||||
|  | } | ||||||
|   | |||||||
							
								
								
									
										39
									
								
								langext/sort.go
									
									
									
									
									
										Normal file
									
								
							
							
						
						
									
										39
									
								
								langext/sort.go
									
									
									
									
									
										Normal file
									
								
							| @@ -0,0 +1,39 @@ | |||||||
|  | package langext | ||||||
|  |  | ||||||
|  | import "sort" | ||||||
|  |  | ||||||
|  | func Sort[T OrderedConstraint](arr []T) { | ||||||
|  | 	sort.Slice(arr, func(i1, i2 int) bool { | ||||||
|  | 		return arr[i1] < arr[i2] | ||||||
|  | 	}) | ||||||
|  | } | ||||||
|  |  | ||||||
|  | func SortStable[T OrderedConstraint](arr []T) { | ||||||
|  | 	sort.SliceStable(arr, func(i1, i2 int) bool { | ||||||
|  | 		return arr[i1] < arr[i2] | ||||||
|  | 	}) | ||||||
|  | } | ||||||
|  |  | ||||||
|  | func IsSorted[T OrderedConstraint](arr []T) bool { | ||||||
|  | 	return sort.SliceIsSorted(arr, func(i1, i2 int) bool { | ||||||
|  | 		return arr[i1] < arr[i2] | ||||||
|  | 	}) | ||||||
|  | } | ||||||
|  |  | ||||||
|  | func SortSlice[T any](arr []T, less func(v1, v2 T) bool) { | ||||||
|  | 	sort.Slice(arr, func(i1, i2 int) bool { | ||||||
|  | 		return less(arr[i1], arr[i2]) | ||||||
|  | 	}) | ||||||
|  | } | ||||||
|  |  | ||||||
|  | func SortSliceStable[T any](arr []T, less func(v1, v2 T) bool) { | ||||||
|  | 	sort.SliceStable(arr, func(i1, i2 int) bool { | ||||||
|  | 		return less(arr[i1], arr[i2]) | ||||||
|  | 	}) | ||||||
|  | } | ||||||
|  |  | ||||||
|  | func IsSliceSorted[T any](arr []T, less func(v1, v2 T) bool) bool { | ||||||
|  | 	return sort.SliceIsSorted(arr, func(i1, i2 int) bool { | ||||||
|  | 		return less(arr[i1], arr[i2]) | ||||||
|  | 	}) | ||||||
|  | } | ||||||
| @@ -107,3 +107,11 @@ func NumToStringOpt[V IntConstraint](v *V, fallback string) string { | |||||||
| 		return fmt.Sprintf("%d", v) | 		return fmt.Sprintf("%d", v) | ||||||
| 	} | 	} | ||||||
| } | } | ||||||
|  |  | ||||||
|  | func StrRepeat(val string, count int) string { | ||||||
|  | 	r := "" | ||||||
|  | 	for i := 0; i < count; i++ { | ||||||
|  | 		r += val | ||||||
|  | 	} | ||||||
|  | 	return r | ||||||
|  | } | ||||||
|   | |||||||
							
								
								
									
										130
									
								
								rext/wrapper.go
									
									
									
									
									
										Normal file
									
								
							
							
						
						
									
										130
									
								
								rext/wrapper.go
									
									
									
									
									
										Normal file
									
								
							| @@ -0,0 +1,130 @@ | |||||||
|  | package rext | ||||||
|  |  | ||||||
|  | import "regexp" | ||||||
|  |  | ||||||
|  | type Regex interface { | ||||||
|  | 	IsMatch(haystack string) bool | ||||||
|  | 	MatchFirst(haystack string) (RegexMatch, bool) | ||||||
|  | 	MatchAll(haystack string) []RegexMatch | ||||||
|  | 	ReplaceAll(haystack string, repl string, literal bool) string | ||||||
|  | 	ReplaceAllFunc(haystack string, repl func(string) string) string | ||||||
|  | 	RemoveAll(haystack string) string | ||||||
|  | 	GroupCount() int | ||||||
|  | } | ||||||
|  |  | ||||||
|  | type regexWrapper struct { | ||||||
|  | 	rex      *regexp.Regexp | ||||||
|  | 	subnames []string | ||||||
|  | } | ||||||
|  |  | ||||||
|  | type RegexMatch struct { | ||||||
|  | 	haystack        string | ||||||
|  | 	submatchesIndex []int | ||||||
|  | 	subnames        []string | ||||||
|  | } | ||||||
|  |  | ||||||
|  | type RegexMatchGroup struct { | ||||||
|  | 	haystack string | ||||||
|  | 	start    int | ||||||
|  | 	end      int | ||||||
|  | } | ||||||
|  |  | ||||||
|  | func W(rex *regexp.Regexp) Regex { | ||||||
|  | 	return ®exWrapper{rex: rex, subnames: rex.SubexpNames()} | ||||||
|  | } | ||||||
|  |  | ||||||
|  | // --------------------------------------------------------------------------------------------------------------------- | ||||||
|  |  | ||||||
|  | func (w *regexWrapper) IsMatch(haystack string) bool { | ||||||
|  | 	return w.rex.MatchString(haystack) | ||||||
|  | } | ||||||
|  |  | ||||||
|  | func (w *regexWrapper) MatchFirst(haystack string) (RegexMatch, bool) { | ||||||
|  | 	res := w.rex.FindStringSubmatchIndex(haystack) | ||||||
|  | 	if res == nil { | ||||||
|  | 		return RegexMatch{}, false | ||||||
|  | 	} | ||||||
|  |  | ||||||
|  | 	return RegexMatch{haystack: haystack, submatchesIndex: res, subnames: w.subnames}, true | ||||||
|  | } | ||||||
|  |  | ||||||
|  | func (w *regexWrapper) MatchAll(haystack string) []RegexMatch { | ||||||
|  | 	resarr := w.rex.FindAllStringSubmatchIndex(haystack, -1) | ||||||
|  |  | ||||||
|  | 	matches := make([]RegexMatch, 0, len(resarr)) | ||||||
|  | 	for _, res := range resarr { | ||||||
|  | 		matches = append(matches, RegexMatch{haystack: haystack, submatchesIndex: res, subnames: w.subnames}) | ||||||
|  | 	} | ||||||
|  |  | ||||||
|  | 	return matches | ||||||
|  | } | ||||||
|  |  | ||||||
|  | func (w *regexWrapper) ReplaceAll(haystack string, repl string, literal bool) string { | ||||||
|  | 	if literal { | ||||||
|  | 		// do not expand placeholder aka $1, $2, ... | ||||||
|  | 		return w.rex.ReplaceAllLiteralString(haystack, repl) | ||||||
|  | 	} else { | ||||||
|  | 		return w.rex.ReplaceAllString(haystack, repl) | ||||||
|  | 	} | ||||||
|  | } | ||||||
|  |  | ||||||
|  | func (w *regexWrapper) ReplaceAllFunc(haystack string, repl func(string) string) string { | ||||||
|  | 	return w.rex.ReplaceAllStringFunc(haystack, repl) | ||||||
|  | } | ||||||
|  |  | ||||||
|  | func (w *regexWrapper) RemoveAll(haystack string) string { | ||||||
|  | 	return w.rex.ReplaceAllLiteralString(haystack, "") | ||||||
|  | } | ||||||
|  |  | ||||||
|  | // GroupCount returns the amount of groups in this match, does not count group-0 (whole match) | ||||||
|  | func (w *regexWrapper) GroupCount() int { | ||||||
|  | 	return len(w.subnames) - 1 | ||||||
|  | } | ||||||
|  |  | ||||||
|  | // --------------------------------------------------------------------------------------------------------------------- | ||||||
|  |  | ||||||
|  | func (m RegexMatch) FullMatch() RegexMatchGroup { | ||||||
|  | 	return RegexMatchGroup{haystack: m.haystack, start: m.submatchesIndex[0], end: m.submatchesIndex[1]} | ||||||
|  | } | ||||||
|  |  | ||||||
|  | // GroupCount returns the amount of groups in this match, does not count group-0 (whole match) | ||||||
|  | func (m RegexMatch) GroupCount() int { | ||||||
|  | 	return len(m.subnames) - 1 | ||||||
|  | } | ||||||
|  |  | ||||||
|  | // GroupByIndex returns the value of a matched group (group 0 == whole match) | ||||||
|  | func (m RegexMatch) GroupByIndex(idx int) RegexMatchGroup { | ||||||
|  | 	return RegexMatchGroup{haystack: m.haystack, start: m.submatchesIndex[idx*2], end: m.submatchesIndex[idx*2+1]} | ||||||
|  | } | ||||||
|  |  | ||||||
|  | // GroupByName returns the value of a matched group (group 0 == whole match) | ||||||
|  | func (m RegexMatch) GroupByName(name string) RegexMatchGroup { | ||||||
|  | 	for idx, subname := range m.subnames { | ||||||
|  | 		if subname == name { | ||||||
|  | 			return RegexMatchGroup{haystack: m.haystack, start: m.submatchesIndex[idx*2], end: m.submatchesIndex[idx*2+1]} | ||||||
|  | 		} | ||||||
|  | 	} | ||||||
|  | 	panic("failed to find regex-group by name") | ||||||
|  | } | ||||||
|  |  | ||||||
|  | // --------------------------------------------------------------------------------------------------------------------- | ||||||
|  |  | ||||||
|  | func (g RegexMatchGroup) Value() string { | ||||||
|  | 	return g.haystack[g.start:g.end] | ||||||
|  | } | ||||||
|  |  | ||||||
|  | func (g RegexMatchGroup) Start() int { | ||||||
|  | 	return g.start | ||||||
|  | } | ||||||
|  |  | ||||||
|  | func (g RegexMatchGroup) End() int { | ||||||
|  | 	return g.end | ||||||
|  | } | ||||||
|  |  | ||||||
|  | func (g RegexMatchGroup) Range() (int, int) { | ||||||
|  | 	return g.start, g.end | ||||||
|  | } | ||||||
|  |  | ||||||
|  | func (g RegexMatchGroup) Length() int { | ||||||
|  | 	return g.end - g.start | ||||||
|  | } | ||||||
							
								
								
									
										128
									
								
								sq/database.go
									
									
									
									
									
										Normal file
									
								
							
							
						
						
									
										128
									
								
								sq/database.go
									
									
									
									
									
										Normal file
									
								
							| @@ -0,0 +1,128 @@ | |||||||
|  | package sq | ||||||
|  |  | ||||||
|  | import ( | ||||||
|  | 	"context" | ||||||
|  | 	"database/sql" | ||||||
|  | 	"github.com/jmoiron/sqlx" | ||||||
|  | 	"sync" | ||||||
|  | ) | ||||||
|  |  | ||||||
|  | type DB interface { | ||||||
|  | 	Exec(ctx context.Context, sql string, prep PP) (sql.Result, error) | ||||||
|  | 	Query(ctx context.Context, sql string, prep PP) (*sqlx.Rows, error) | ||||||
|  | 	Ping(ctx context.Context) error | ||||||
|  | 	BeginTransaction(ctx context.Context, iso sql.IsolationLevel) (Tx, error) | ||||||
|  | 	AddListener(listener Listener) | ||||||
|  | 	Exit() error | ||||||
|  | } | ||||||
|  |  | ||||||
|  | type database struct { | ||||||
|  | 	db    *sqlx.DB | ||||||
|  | 	txctr uint16 | ||||||
|  | 	lock  sync.Mutex | ||||||
|  | 	lstr  []Listener | ||||||
|  | } | ||||||
|  |  | ||||||
|  | func NewDB(db *sqlx.DB) DB { | ||||||
|  | 	return &database{ | ||||||
|  | 		db:    db, | ||||||
|  | 		txctr: 0, | ||||||
|  | 		lock:  sync.Mutex{}, | ||||||
|  | 		lstr:  make([]Listener, 0), | ||||||
|  | 	} | ||||||
|  | } | ||||||
|  |  | ||||||
|  | func (db *database) AddListener(listener Listener) { | ||||||
|  | 	db.lstr = append(db.lstr, listener) | ||||||
|  | } | ||||||
|  |  | ||||||
|  | func (db *database) Exec(ctx context.Context, sqlstr string, prep PP) (sql.Result, error) { | ||||||
|  | 	origsql := sqlstr | ||||||
|  | 	for _, v := range db.lstr { | ||||||
|  | 		err := v.PreExec(ctx, nil, &sqlstr, &prep) | ||||||
|  | 		if err != nil { | ||||||
|  | 			return nil, err | ||||||
|  | 		} | ||||||
|  | 	} | ||||||
|  |  | ||||||
|  | 	res, err := db.db.NamedExecContext(ctx, sqlstr, prep) | ||||||
|  |  | ||||||
|  | 	for _, v := range db.lstr { | ||||||
|  | 		v.PostExec(nil, origsql, sqlstr, prep) | ||||||
|  | 	} | ||||||
|  |  | ||||||
|  | 	if err != nil { | ||||||
|  | 		return nil, err | ||||||
|  | 	} | ||||||
|  | 	return res, nil | ||||||
|  | } | ||||||
|  |  | ||||||
|  | func (db *database) Query(ctx context.Context, sqlstr string, prep PP) (*sqlx.Rows, error) { | ||||||
|  | 	origsql := sqlstr | ||||||
|  | 	for _, v := range db.lstr { | ||||||
|  | 		err := v.PreQuery(ctx, nil, &sqlstr, &prep) | ||||||
|  | 		if err != nil { | ||||||
|  | 			return nil, err | ||||||
|  | 		} | ||||||
|  | 	} | ||||||
|  |  | ||||||
|  | 	rows, err := sqlx.NamedQueryContext(ctx, db.db, sqlstr, prep) | ||||||
|  |  | ||||||
|  | 	for _, v := range db.lstr { | ||||||
|  | 		v.PostQuery(nil, origsql, sqlstr, prep) | ||||||
|  | 	} | ||||||
|  |  | ||||||
|  | 	if err != nil { | ||||||
|  | 		return nil, err | ||||||
|  | 	} | ||||||
|  | 	return rows, nil | ||||||
|  | } | ||||||
|  |  | ||||||
|  | func (db *database) Ping(ctx context.Context) error { | ||||||
|  | 	for _, v := range db.lstr { | ||||||
|  | 		err := v.PrePing(ctx) | ||||||
|  | 		if err != nil { | ||||||
|  | 			return err | ||||||
|  | 		} | ||||||
|  | 	} | ||||||
|  |  | ||||||
|  | 	err := db.db.PingContext(ctx) | ||||||
|  |  | ||||||
|  | 	for _, v := range db.lstr { | ||||||
|  | 		v.PostPing(err) | ||||||
|  | 	} | ||||||
|  |  | ||||||
|  | 	if err != nil { | ||||||
|  | 		return err | ||||||
|  | 	} | ||||||
|  | 	return nil | ||||||
|  | } | ||||||
|  |  | ||||||
|  | func (db *database) BeginTransaction(ctx context.Context, iso sql.IsolationLevel) (Tx, error) { | ||||||
|  | 	db.lock.Lock() | ||||||
|  | 	txid := db.txctr | ||||||
|  | 	db.txctr += 1 // with overflow ! | ||||||
|  | 	db.lock.Unlock() | ||||||
|  |  | ||||||
|  | 	for _, v := range db.lstr { | ||||||
|  | 		err := v.PreTxBegin(ctx, txid) | ||||||
|  | 		if err != nil { | ||||||
|  | 			return nil, err | ||||||
|  | 		} | ||||||
|  | 	} | ||||||
|  |  | ||||||
|  | 	xtx, err := db.db.BeginTxx(ctx, &sql.TxOptions{Isolation: iso}) | ||||||
|  | 	if err != nil { | ||||||
|  | 		return nil, err | ||||||
|  | 	} | ||||||
|  |  | ||||||
|  | 	for _, v := range db.lstr { | ||||||
|  | 		v.PostTxBegin(txid, err) | ||||||
|  | 	} | ||||||
|  |  | ||||||
|  | 	return NewTransaction(xtx, txid, db.lstr), nil | ||||||
|  | } | ||||||
|  |  | ||||||
|  | func (db *database) Exit() error { | ||||||
|  | 	return db.db.Close() | ||||||
|  | } | ||||||
							
								
								
									
										19
									
								
								sq/listener.go
									
									
									
									
									
										Normal file
									
								
							
							
						
						
									
										19
									
								
								sq/listener.go
									
									
									
									
									
										Normal file
									
								
							| @@ -0,0 +1,19 @@ | |||||||
|  | package sq | ||||||
|  |  | ||||||
|  | import "context" | ||||||
|  |  | ||||||
|  | type Listener interface { | ||||||
|  | 	PrePing(ctx context.Context) error | ||||||
|  | 	PreTxBegin(ctx context.Context, txid uint16) error | ||||||
|  | 	PreTxCommit(txid uint16) error | ||||||
|  | 	PreTxRollback(txid uint16) error | ||||||
|  | 	PreQuery(ctx context.Context, txID *uint16, sql *string, params *PP) error | ||||||
|  | 	PreExec(ctx context.Context, txID *uint16, sql *string, params *PP) error | ||||||
|  |  | ||||||
|  | 	PostPing(result error) | ||||||
|  | 	PostTxBegin(txid uint16, result error) | ||||||
|  | 	PostTxCommit(txid uint16, result error) | ||||||
|  | 	PostTxRollback(txid uint16, result error) | ||||||
|  | 	PostQuery(txID *uint16, sqlOriginal string, sqlReal string, params PP) | ||||||
|  | 	PostExec(txID *uint16, sqlOriginal string, sqlReal string, params PP) | ||||||
|  | } | ||||||
							
								
								
									
										13
									
								
								sq/params.go
									
									
									
									
									
										Normal file
									
								
							
							
						
						
									
										13
									
								
								sq/params.go
									
									
									
									
									
										Normal file
									
								
							| @@ -0,0 +1,13 @@ | |||||||
|  | package sq | ||||||
|  |  | ||||||
|  | type PP map[string]any | ||||||
|  |  | ||||||
|  | func Join(pps ...PP) PP { | ||||||
|  | 	r := PP{} | ||||||
|  | 	for _, add := range pps { | ||||||
|  | 		for k, v := range add { | ||||||
|  | 			r[k] = v | ||||||
|  | 		} | ||||||
|  | 	} | ||||||
|  | 	return r | ||||||
|  | } | ||||||
							
								
								
									
										12
									
								
								sq/queryable.go
									
									
									
									
									
										Normal file
									
								
							
							
						
						
									
										12
									
								
								sq/queryable.go
									
									
									
									
									
										Normal file
									
								
							| @@ -0,0 +1,12 @@ | |||||||
|  | package sq | ||||||
|  |  | ||||||
|  | import ( | ||||||
|  | 	"context" | ||||||
|  | 	"database/sql" | ||||||
|  | 	"github.com/jmoiron/sqlx" | ||||||
|  | ) | ||||||
|  |  | ||||||
|  | type Queryable interface { | ||||||
|  | 	Exec(ctx context.Context, sql string, prep PP) (sql.Result, error) | ||||||
|  | 	Query(ctx context.Context, sql string, prep PP) (*sqlx.Rows, error) | ||||||
|  | } | ||||||
							
								
								
									
										140
									
								
								sq/scanner.go
									
									
									
									
									
										Normal file
									
								
							
							
						
						
									
										140
									
								
								sq/scanner.go
									
									
									
									
									
										Normal file
									
								
							| @@ -0,0 +1,140 @@ | |||||||
|  | package sq | ||||||
|  |  | ||||||
|  | import ( | ||||||
|  | 	"database/sql" | ||||||
|  | 	"errors" | ||||||
|  | 	"github.com/jmoiron/sqlx" | ||||||
|  | ) | ||||||
|  |  | ||||||
|  | type StructScanMode string | ||||||
|  |  | ||||||
|  | const ( | ||||||
|  | 	SModeFast     StructScanMode = "FAST" | ||||||
|  | 	SModeExtended StructScanMode = "EXTENDED" | ||||||
|  | ) | ||||||
|  |  | ||||||
|  | type StructScanSafety string | ||||||
|  |  | ||||||
|  | const ( | ||||||
|  | 	Safe   StructScanSafety = "SAFE" | ||||||
|  | 	Unsafe StructScanSafety = "UNSAFE" | ||||||
|  | ) | ||||||
|  |  | ||||||
|  | func ScanSingle[TData any](rows *sqlx.Rows, mode StructScanMode, sec StructScanSafety, close bool) (TData, error) { | ||||||
|  | 	if rows.Next() { | ||||||
|  | 		var strscan *StructScanner | ||||||
|  |  | ||||||
|  | 		if sec == Safe { | ||||||
|  | 			strscan = NewStructScanner(rows, false) | ||||||
|  | 			var data TData | ||||||
|  | 			err := strscan.Start(&data) | ||||||
|  | 			if err != nil { | ||||||
|  | 				return *new(TData), err | ||||||
|  | 			} | ||||||
|  | 		} else if sec == Unsafe { | ||||||
|  | 			strscan = NewStructScanner(rows, true) | ||||||
|  | 			var data TData | ||||||
|  | 			err := strscan.Start(&data) | ||||||
|  | 			if err != nil { | ||||||
|  | 				return *new(TData), err | ||||||
|  | 			} | ||||||
|  | 		} else { | ||||||
|  | 			return *new(TData), errors.New("unknown value for <sec>") | ||||||
|  | 		} | ||||||
|  |  | ||||||
|  | 		var data TData | ||||||
|  |  | ||||||
|  | 		if mode == SModeFast { | ||||||
|  | 			err := strscan.StructScanBase(&data) | ||||||
|  | 			if err != nil { | ||||||
|  | 				return *new(TData), err | ||||||
|  | 			} | ||||||
|  | 		} else if mode == SModeExtended { | ||||||
|  | 			err := strscan.StructScanExt(&data) | ||||||
|  | 			if err != nil { | ||||||
|  | 				return *new(TData), err | ||||||
|  | 			} | ||||||
|  | 		} else { | ||||||
|  | 			return *new(TData), errors.New("unknown value for <mode>") | ||||||
|  | 		} | ||||||
|  |  | ||||||
|  | 		if rows.Next() { | ||||||
|  | 			if close { | ||||||
|  | 				_ = rows.Close() | ||||||
|  | 			} | ||||||
|  | 			return *new(TData), errors.New("sql returned more than one row") | ||||||
|  | 		} | ||||||
|  |  | ||||||
|  | 		if close { | ||||||
|  | 			err := rows.Close() | ||||||
|  | 			if err != nil { | ||||||
|  | 				return *new(TData), err | ||||||
|  | 			} | ||||||
|  | 		} | ||||||
|  |  | ||||||
|  | 		if err := rows.Err(); err != nil { | ||||||
|  | 			return *new(TData), err | ||||||
|  | 		} | ||||||
|  |  | ||||||
|  | 		return data, nil | ||||||
|  |  | ||||||
|  | 	} else { | ||||||
|  | 		if close { | ||||||
|  | 			_ = rows.Close() | ||||||
|  | 		} | ||||||
|  | 		return *new(TData), sql.ErrNoRows | ||||||
|  | 	} | ||||||
|  | } | ||||||
|  |  | ||||||
|  | func ScanAll[TData any](rows *sqlx.Rows, mode StructScanMode, sec StructScanSafety, close bool) ([]TData, error) { | ||||||
|  | 	var strscan *StructScanner | ||||||
|  |  | ||||||
|  | 	if sec == Safe { | ||||||
|  | 		strscan = NewStructScanner(rows, false) | ||||||
|  | 		var data TData | ||||||
|  | 		err := strscan.Start(&data) | ||||||
|  | 		if err != nil { | ||||||
|  | 			return nil, err | ||||||
|  | 		} | ||||||
|  | 	} else if sec == Unsafe { | ||||||
|  | 		strscan = NewStructScanner(rows, true) | ||||||
|  | 		var data TData | ||||||
|  | 		err := strscan.Start(&data) | ||||||
|  | 		if err != nil { | ||||||
|  | 			return nil, err | ||||||
|  | 		} | ||||||
|  | 	} else { | ||||||
|  | 		return nil, errors.New("unknown value for <sec>") | ||||||
|  | 	} | ||||||
|  |  | ||||||
|  | 	res := make([]TData, 0) | ||||||
|  | 	for rows.Next() { | ||||||
|  | 		if mode == SModeFast { | ||||||
|  | 			var data TData | ||||||
|  | 			err := strscan.StructScanBase(&data) | ||||||
|  | 			if err != nil { | ||||||
|  | 				return nil, err | ||||||
|  | 			} | ||||||
|  | 			res = append(res, data) | ||||||
|  | 		} else if mode == SModeExtended { | ||||||
|  | 			var data TData | ||||||
|  | 			err := strscan.StructScanExt(&data) | ||||||
|  | 			if err != nil { | ||||||
|  | 				return nil, err | ||||||
|  | 			} | ||||||
|  | 			res = append(res, data) | ||||||
|  | 		} else { | ||||||
|  | 			return nil, errors.New("unknown value for <mode>") | ||||||
|  | 		} | ||||||
|  | 	} | ||||||
|  | 	if close { | ||||||
|  | 		err := strscan.rows.Close() | ||||||
|  | 		if err != nil { | ||||||
|  | 			return nil, err | ||||||
|  | 		} | ||||||
|  | 	} | ||||||
|  | 	if err := rows.Err(); err != nil { | ||||||
|  | 		return nil, err | ||||||
|  | 	} | ||||||
|  | 	return res, nil | ||||||
|  | } | ||||||
							
								
								
									
										223
									
								
								sq/structscanner.go
									
									
									
									
									
										Normal file
									
								
							
							
						
						
									
										223
									
								
								sq/structscanner.go
									
									
									
									
									
										Normal file
									
								
							| @@ -0,0 +1,223 @@ | |||||||
|  | package sq | ||||||
|  |  | ||||||
|  | import ( | ||||||
|  | 	"errors" | ||||||
|  | 	"fmt" | ||||||
|  | 	"github.com/jmoiron/sqlx" | ||||||
|  | 	"github.com/jmoiron/sqlx/reflectx" | ||||||
|  | 	"reflect" | ||||||
|  | ) | ||||||
|  |  | ||||||
|  | // forked from sqlx, but added ability to unmarshal optional-nested structs | ||||||
|  |  | ||||||
|  | type StructScanner struct { | ||||||
|  | 	rows   *sqlx.Rows | ||||||
|  | 	Mapper *reflectx.Mapper | ||||||
|  | 	unsafe bool | ||||||
|  |  | ||||||
|  | 	fields  [][]int | ||||||
|  | 	values  []any | ||||||
|  | 	columns []string | ||||||
|  | } | ||||||
|  |  | ||||||
|  | func NewStructScanner(rows *sqlx.Rows, unsafe bool) *StructScanner { | ||||||
|  | 	return &StructScanner{ | ||||||
|  | 		rows:   rows, | ||||||
|  | 		Mapper: reflectx.NewMapper("db"), | ||||||
|  | 		unsafe: unsafe, | ||||||
|  | 	} | ||||||
|  | } | ||||||
|  |  | ||||||
|  | func (r *StructScanner) Start(dest any) error { | ||||||
|  | 	v := reflect.ValueOf(dest) | ||||||
|  |  | ||||||
|  | 	if v.Kind() != reflect.Ptr { | ||||||
|  | 		return errors.New("must pass a pointer, not a value, to StructScan destination") | ||||||
|  | 	} | ||||||
|  |  | ||||||
|  | 	columns, err := r.rows.Columns() | ||||||
|  | 	if err != nil { | ||||||
|  | 		return err | ||||||
|  | 	} | ||||||
|  |  | ||||||
|  | 	r.columns = columns | ||||||
|  | 	r.fields = r.Mapper.TraversalsByName(v.Type(), columns) | ||||||
|  | 	// if we are not unsafe and are missing fields, return an error | ||||||
|  | 	if f, err := missingFields(r.fields); err != nil && !r.unsafe { | ||||||
|  | 		return fmt.Errorf("missing destination name %s in %T", columns[f], dest) | ||||||
|  | 	} | ||||||
|  | 	r.values = make([]interface{}, len(columns)) | ||||||
|  |  | ||||||
|  | 	return nil | ||||||
|  | } | ||||||
|  |  | ||||||
|  | // StructScanExt forked from github.com/jmoiron/sqlx@v1.3.5/sqlx.go | ||||||
|  | // does also wok with nullabel structs (from LEFT JOIN's) | ||||||
|  | func (r *StructScanner) StructScanExt(dest any) error { | ||||||
|  | 	v := reflect.ValueOf(dest) | ||||||
|  |  | ||||||
|  | 	if v.Kind() != reflect.Ptr { | ||||||
|  | 		return errors.New("must pass a pointer, not a value, to StructScan destination") | ||||||
|  | 	} | ||||||
|  |  | ||||||
|  | 	// ========= STEP 1 ::  ========= | ||||||
|  |  | ||||||
|  | 	v = v.Elem() | ||||||
|  |  | ||||||
|  | 	err := fieldsByTraversalExtended(v, r.fields, r.values) | ||||||
|  | 	if err != nil { | ||||||
|  | 		return err | ||||||
|  | 	} | ||||||
|  | 	// scan into the struct field pointers and append to our results | ||||||
|  | 	err = r.rows.Scan(r.values...) | ||||||
|  | 	if err != nil { | ||||||
|  | 		return err | ||||||
|  | 	} | ||||||
|  |  | ||||||
|  | 	nullStructs := make(map[string]bool) | ||||||
|  |  | ||||||
|  | 	for i, traversal := range r.fields { | ||||||
|  | 		if len(traversal) == 0 { | ||||||
|  | 			continue | ||||||
|  | 		} | ||||||
|  |  | ||||||
|  | 		isnsil := reflect.ValueOf(r.values[i]).Elem().IsNil() | ||||||
|  |  | ||||||
|  | 		for i := 1; i < len(traversal); i++ { | ||||||
|  |  | ||||||
|  | 			canParentNil := reflectx.FieldByIndexes(v, traversal[0:i]).Kind() == reflect.Pointer | ||||||
|  |  | ||||||
|  | 			k := fmt.Sprintf("%v", traversal[0:i]) | ||||||
|  | 			if v, ok := nullStructs[k]; ok { | ||||||
|  |  | ||||||
|  | 				nullStructs[k] = canParentNil && v && isnsil | ||||||
|  |  | ||||||
|  | 			} else { | ||||||
|  | 				nullStructs[k] = canParentNil && isnsil | ||||||
|  | 			} | ||||||
|  | 		} | ||||||
|  |  | ||||||
|  | 	} | ||||||
|  |  | ||||||
|  | 	forcenulled := make(map[string]bool) | ||||||
|  |  | ||||||
|  | 	for i, traversal := range r.fields { | ||||||
|  | 		if len(traversal) == 0 { | ||||||
|  | 			continue | ||||||
|  | 		} | ||||||
|  |  | ||||||
|  | 		anyparentnull := false | ||||||
|  | 		for i := 1; i < len(traversal); i++ { | ||||||
|  | 			k := fmt.Sprintf("%v", traversal[0:i]) | ||||||
|  | 			if nv, ok := nullStructs[k]; ok && nv { | ||||||
|  |  | ||||||
|  | 				if _, ok := forcenulled[k]; !ok { | ||||||
|  | 					f := reflectx.FieldByIndexes(v, traversal[0:i]) | ||||||
|  | 					f.Set(reflect.Zero(f.Type())) // set to nil | ||||||
|  | 					forcenulled[k] = true | ||||||
|  | 				} | ||||||
|  |  | ||||||
|  | 				anyparentnull = true | ||||||
|  | 				break | ||||||
|  |  | ||||||
|  | 			} | ||||||
|  | 		} | ||||||
|  |  | ||||||
|  | 		if anyparentnull { | ||||||
|  | 			continue | ||||||
|  | 		} | ||||||
|  |  | ||||||
|  | 		f := reflectx.FieldByIndexes(v, traversal) | ||||||
|  |  | ||||||
|  | 		val1 := reflect.ValueOf(r.values[i]) | ||||||
|  | 		val2 := val1.Elem() | ||||||
|  | 		val3 := val2.Elem() | ||||||
|  |  | ||||||
|  | 		if val2.IsNil() { | ||||||
|  | 			if f.Kind() != reflect.Pointer { | ||||||
|  | 				return errors.New(fmt.Sprintf("Cannot set field %v to NULL value from column '%s' (type: %s)", traversal, r.columns[i], f.Type().String())) | ||||||
|  | 			} | ||||||
|  |  | ||||||
|  | 			f.Set(reflect.Zero(f.Type())) // set to nil | ||||||
|  | 		} else { | ||||||
|  | 			f.Set(val3) | ||||||
|  | 		} | ||||||
|  |  | ||||||
|  | 	} | ||||||
|  |  | ||||||
|  | 	return r.rows.Err() | ||||||
|  | } | ||||||
|  |  | ||||||
|  | // StructScanBase forked from github.com/jmoiron/sqlx@v1.3.5/sqlx.go | ||||||
|  | // without (relevant) changes | ||||||
|  | func (r *StructScanner) StructScanBase(dest any) error { | ||||||
|  | 	v := reflect.ValueOf(dest) | ||||||
|  |  | ||||||
|  | 	if v.Kind() != reflect.Ptr { | ||||||
|  | 		return errors.New("must pass a pointer, not a value, to StructScan destination") | ||||||
|  | 	} | ||||||
|  |  | ||||||
|  | 	v = v.Elem() | ||||||
|  |  | ||||||
|  | 	err := fieldsByTraversalBase(v, r.fields, r.values, true) | ||||||
|  | 	if err != nil { | ||||||
|  | 		return err | ||||||
|  | 	} | ||||||
|  | 	// scan into the struct field pointers and append to our results | ||||||
|  | 	err = r.rows.Scan(r.values...) | ||||||
|  | 	if err != nil { | ||||||
|  | 		return err | ||||||
|  | 	} | ||||||
|  | 	return r.rows.Err() | ||||||
|  | } | ||||||
|  |  | ||||||
|  | // fieldsByTraversal forked from github.com/jmoiron/sqlx@v1.3.5/sqlx.go | ||||||
|  | func fieldsByTraversalExtended(v reflect.Value, traversals [][]int, values []interface{}) error { | ||||||
|  | 	v = reflect.Indirect(v) | ||||||
|  | 	if v.Kind() != reflect.Struct { | ||||||
|  | 		return errors.New("argument not a struct") | ||||||
|  | 	} | ||||||
|  |  | ||||||
|  | 	for i, traversal := range traversals { | ||||||
|  | 		if len(traversal) == 0 { | ||||||
|  | 			values[i] = new(interface{}) | ||||||
|  | 			continue | ||||||
|  | 		} | ||||||
|  | 		f := reflectx.FieldByIndexes(v, traversal) | ||||||
|  |  | ||||||
|  | 		values[i] = reflect.New(reflect.PointerTo(f.Type())).Interface() | ||||||
|  | 	} | ||||||
|  | 	return nil | ||||||
|  | } | ||||||
|  |  | ||||||
|  | // fieldsByTraversal forked from github.com/jmoiron/sqlx@v1.3.5/sqlx.go | ||||||
|  | func fieldsByTraversalBase(v reflect.Value, traversals [][]int, values []interface{}, ptrs bool) error { | ||||||
|  | 	v = reflect.Indirect(v) | ||||||
|  | 	if v.Kind() != reflect.Struct { | ||||||
|  | 		return errors.New("argument not a struct") | ||||||
|  | 	} | ||||||
|  |  | ||||||
|  | 	for i, traversal := range traversals { | ||||||
|  | 		if len(traversal) == 0 { | ||||||
|  | 			values[i] = new(interface{}) | ||||||
|  | 			continue | ||||||
|  | 		} | ||||||
|  | 		f := reflectx.FieldByIndexes(v, traversal) | ||||||
|  | 		if ptrs { | ||||||
|  | 			values[i] = f.Addr().Interface() | ||||||
|  | 		} else { | ||||||
|  | 			values[i] = f.Interface() | ||||||
|  | 		} | ||||||
|  | 	} | ||||||
|  | 	return nil | ||||||
|  | } | ||||||
|  |  | ||||||
|  | // missingFields forked from github.com/jmoiron/sqlx@v1.3.5/sqlx.go | ||||||
|  | func missingFields(transversals [][]int) (field int, err error) { | ||||||
|  | 	for i, t := range transversals { | ||||||
|  | 		if len(t) == 0 { | ||||||
|  | 			return i, errors.New("missing field") | ||||||
|  | 		} | ||||||
|  | 	} | ||||||
|  | 	return 0, nil | ||||||
|  | } | ||||||
							
								
								
									
										105
									
								
								sq/transaction.go
									
									
									
									
									
										Normal file
									
								
							
							
						
						
									
										105
									
								
								sq/transaction.go
									
									
									
									
									
										Normal file
									
								
							| @@ -0,0 +1,105 @@ | |||||||
|  | package sq | ||||||
|  |  | ||||||
|  | import ( | ||||||
|  | 	"context" | ||||||
|  | 	"database/sql" | ||||||
|  | 	"github.com/jmoiron/sqlx" | ||||||
|  | 	"gogs.mikescher.com/BlackForestBytes/goext/langext" | ||||||
|  | ) | ||||||
|  |  | ||||||
|  | type Tx interface { | ||||||
|  | 	Rollback() error | ||||||
|  | 	Commit() error | ||||||
|  | 	Exec(ctx context.Context, sql string, prep PP) (sql.Result, error) | ||||||
|  | 	Query(ctx context.Context, sql string, prep PP) (*sqlx.Rows, error) | ||||||
|  | } | ||||||
|  |  | ||||||
|  | type transaction struct { | ||||||
|  | 	tx   *sqlx.Tx | ||||||
|  | 	id   uint16 | ||||||
|  | 	lstr []Listener | ||||||
|  | } | ||||||
|  |  | ||||||
|  | func NewTransaction(xtx *sqlx.Tx, txid uint16, lstr []Listener) Tx { | ||||||
|  | 	return &transaction{ | ||||||
|  | 		tx:   xtx, | ||||||
|  | 		id:   txid, | ||||||
|  | 		lstr: lstr, | ||||||
|  | 	} | ||||||
|  | } | ||||||
|  |  | ||||||
|  | func (tx *transaction) Rollback() error { | ||||||
|  | 	for _, v := range tx.lstr { | ||||||
|  | 		err := v.PreTxRollback(tx.id) | ||||||
|  | 		if err != nil { | ||||||
|  | 			return err | ||||||
|  | 		} | ||||||
|  | 	} | ||||||
|  |  | ||||||
|  | 	result := tx.tx.Rollback() | ||||||
|  |  | ||||||
|  | 	for _, v := range tx.lstr { | ||||||
|  | 		v.PostTxRollback(tx.id, result) | ||||||
|  | 	} | ||||||
|  |  | ||||||
|  | 	return result | ||||||
|  | } | ||||||
|  |  | ||||||
|  | func (tx *transaction) Commit() error { | ||||||
|  | 	for _, v := range tx.lstr { | ||||||
|  | 		err := v.PreTxCommit(tx.id) | ||||||
|  | 		if err != nil { | ||||||
|  | 			return err | ||||||
|  | 		} | ||||||
|  | 	} | ||||||
|  |  | ||||||
|  | 	result := tx.tx.Commit() | ||||||
|  |  | ||||||
|  | 	for _, v := range tx.lstr { | ||||||
|  | 		v.PostTxRollback(tx.id, result) | ||||||
|  | 	} | ||||||
|  |  | ||||||
|  | 	return result | ||||||
|  | } | ||||||
|  |  | ||||||
|  | func (tx *transaction) Exec(ctx context.Context, sqlstr string, prep PP) (sql.Result, error) { | ||||||
|  | 	origsql := sqlstr | ||||||
|  | 	for _, v := range tx.lstr { | ||||||
|  | 		err := v.PreExec(ctx, langext.Ptr(tx.id), &sqlstr, &prep) | ||||||
|  | 		if err != nil { | ||||||
|  | 			return nil, err | ||||||
|  | 		} | ||||||
|  | 	} | ||||||
|  |  | ||||||
|  | 	res, err := tx.tx.NamedExecContext(ctx, sqlstr, prep) | ||||||
|  |  | ||||||
|  | 	for _, v := range tx.lstr { | ||||||
|  | 		v.PostExec(langext.Ptr(tx.id), origsql, sqlstr, prep) | ||||||
|  | 	} | ||||||
|  |  | ||||||
|  | 	if err != nil { | ||||||
|  | 		return nil, err | ||||||
|  | 	} | ||||||
|  | 	return res, nil | ||||||
|  | } | ||||||
|  |  | ||||||
|  | func (tx *transaction) Query(ctx context.Context, sqlstr string, prep PP) (*sqlx.Rows, error) { | ||||||
|  | 	origsql := sqlstr | ||||||
|  | 	for _, v := range tx.lstr { | ||||||
|  | 		err := v.PreQuery(ctx, langext.Ptr(tx.id), &sqlstr, &prep) | ||||||
|  | 		if err != nil { | ||||||
|  | 			return nil, err | ||||||
|  | 		} | ||||||
|  | 	} | ||||||
|  |  | ||||||
|  | 	rows, err := sqlx.NamedQueryContext(ctx, tx.tx, sqlstr, prep) | ||||||
|  |  | ||||||
|  | 	for _, v := range tx.lstr { | ||||||
|  | 		v.PostQuery(langext.Ptr(tx.id), origsql, sqlstr, prep) | ||||||
|  | 	} | ||||||
|  |  | ||||||
|  | 	if err != nil { | ||||||
|  | 		return nil, err | ||||||
|  | 	} | ||||||
|  | 	return rows, nil | ||||||
|  | } | ||||||
		Reference in New Issue
	
	Block a user