Browse Source
- 整包删除(只有自身测试引用):lego/sys/{gin,timewheel,lghttp,sdk/bytedance/tos}、
lego/utils/crypto/{gm,gm_java,sra,base64}、lego/utils/container 根目录及 addr/ip/sortslice/version
(container/id 保留)、sys/{dify,deepseek,qweather,coze,axml,haifanwu,aliyun/sts,websearch/bingsearch,websearch/bravesearch}
- mcp 的 tool_finance / tool_spotify_music_* 从未 RegisterComp,连带 sys/juhe、sys/spotify 一起删;
sys/openai 测试里依赖 juhe 的 finance 工具用例随之去掉
- modules/timer/timer_uselog.go 从未装上;comm/const.go 四个零引用的 Module* 常量
- 配置模板去掉没人读的段:home 的 deepseek/ali_filetrans、mcp 的 juhe/spotify/bravesearch、
console 的 cos、gateway 白名单里没有对应路由的 public_info
验证:go build/vet 通过;go test 除两个本来就依赖本机 MySQL / 外部 HTTP 的集成测试
(lego/sys/mysql、modules/agents,改前同样失败)外全部通过。
Co-Authored-By: Claude Fable 5.1 <noreply@anthropic.com>
(cherry picked from commit 990c09f066e63a31c483adff1a6d2e3c9ab72c3e)
main
129 changed files with 2 additions and 11748 deletions
@ -1,73 +0,0 @@ |
|||||
package binding |
|
||||
|
|
||||
import "net/http" |
|
||||
|
|
||||
// Content-Type MIME of the most common data formats.
|
|
||||
const ( |
|
||||
MIMEJSON = "application/json" |
|
||||
MIMEHTML = "text/html" |
|
||||
MIMEXML = "application/xml" |
|
||||
MIMEXML2 = "text/xml" |
|
||||
MIMEPlain = "text/plain" |
|
||||
MIMEPOSTForm = "application/x-www-form-urlencoded" |
|
||||
MIMEMultipartPOSTForm = "multipart/form-data" |
|
||||
MIMEPROTOBUF = "application/x-protobuf" |
|
||||
MIMEMSGPACK = "application/x-msgpack" |
|
||||
MIMEMSGPACK2 = "application/msgpack" |
|
||||
MIMEYAML = "application/x-yaml" |
|
||||
) |
|
||||
|
|
||||
type Binding interface { |
|
||||
Name() string |
|
||||
Bind(*http.Request, interface{}) error |
|
||||
} |
|
||||
|
|
||||
type StructValidator interface { |
|
||||
ValidateStruct(interface{}) error |
|
||||
Engine() interface{} |
|
||||
} |
|
||||
|
|
||||
var Validator StructValidator = &defaultValidator{} |
|
||||
var ( |
|
||||
JSON = jsonBinding{} |
|
||||
XML = xmlBinding{} |
|
||||
Form = formBinding{} |
|
||||
Query = queryBinding{} |
|
||||
FormPost = formPostBinding{} |
|
||||
FormMultipart = formMultipartBinding{} |
|
||||
ProtoBuf = protobufBinding{} |
|
||||
MsgPack = msgpackBinding{} |
|
||||
YAML = yamlBinding{} |
|
||||
Uri = uriBinding{} |
|
||||
Header = headerBinding{} |
|
||||
) |
|
||||
|
|
||||
func Default(method, contentType string) Binding { |
|
||||
if method == http.MethodGet { |
|
||||
return Form |
|
||||
} |
|
||||
|
|
||||
switch contentType { |
|
||||
case MIMEJSON: |
|
||||
return JSON |
|
||||
case MIMEXML, MIMEXML2: |
|
||||
return XML |
|
||||
case MIMEPROTOBUF: |
|
||||
return ProtoBuf |
|
||||
case MIMEMSGPACK, MIMEMSGPACK2: |
|
||||
return MsgPack |
|
||||
case MIMEYAML: |
|
||||
return YAML |
|
||||
case MIMEMultipartPOSTForm: |
|
||||
return FormMultipart |
|
||||
default: // case MIMEPOSTForm:
|
|
||||
return Form |
|
||||
} |
|
||||
} |
|
||||
|
|
||||
func validate(obj interface{}) error { |
|
||||
if Validator == nil { |
|
||||
return nil |
|
||||
} |
|
||||
return Validator.ValidateStruct(obj) |
|
||||
} |
|
||||
@ -1,97 +0,0 @@ |
|||||
// Copyright 2017 Manu Martinez-Almeida. All rights reserved.
|
|
||||
// Use of this source code is governed by a MIT style
|
|
||||
// license that can be found in the LICENSE file.
|
|
||||
|
|
||||
package binding |
|
||||
|
|
||||
import ( |
|
||||
"fmt" |
|
||||
"reflect" |
|
||||
"strings" |
|
||||
"sync" |
|
||||
|
|
||||
"github.com/go-playground/validator/v10" |
|
||||
) |
|
||||
|
|
||||
type defaultValidator struct { |
|
||||
once sync.Once |
|
||||
validate *validator.Validate |
|
||||
} |
|
||||
|
|
||||
type SliceValidationError []error |
|
||||
|
|
||||
// Error concatenates all error elements in SliceValidationError into a single string separated by \n.
|
|
||||
func (err SliceValidationError) Error() string { |
|
||||
n := len(err) |
|
||||
switch n { |
|
||||
case 0: |
|
||||
return "" |
|
||||
default: |
|
||||
var b strings.Builder |
|
||||
if err[0] != nil { |
|
||||
fmt.Fprintf(&b, "[%d]: %s", 0, err[0].Error()) |
|
||||
} |
|
||||
if n > 1 { |
|
||||
for i := 1; i < n; i++ { |
|
||||
if err[i] != nil { |
|
||||
b.WriteString("\n") |
|
||||
fmt.Fprintf(&b, "[%d]: %s", i, err[i].Error()) |
|
||||
} |
|
||||
} |
|
||||
} |
|
||||
return b.String() |
|
||||
} |
|
||||
} |
|
||||
|
|
||||
var _ StructValidator = &defaultValidator{} |
|
||||
|
|
||||
// ValidateStruct receives any kind of type, but only performed struct or pointer to struct type.
|
|
||||
func (v *defaultValidator) ValidateStruct(obj interface{}) error { |
|
||||
if obj == nil { |
|
||||
return nil |
|
||||
} |
|
||||
|
|
||||
value := reflect.ValueOf(obj) |
|
||||
switch value.Kind() { |
|
||||
case reflect.Ptr: |
|
||||
return v.ValidateStruct(value.Elem().Interface()) |
|
||||
case reflect.Struct: |
|
||||
return v.validateStruct(obj) |
|
||||
case reflect.Slice, reflect.Array: |
|
||||
count := value.Len() |
|
||||
validateRet := make(SliceValidationError, 0) |
|
||||
for i := 0; i < count; i++ { |
|
||||
if err := v.ValidateStruct(value.Index(i).Interface()); err != nil { |
|
||||
validateRet = append(validateRet, err) |
|
||||
} |
|
||||
} |
|
||||
if len(validateRet) == 0 { |
|
||||
return nil |
|
||||
} |
|
||||
return validateRet |
|
||||
default: |
|
||||
return nil |
|
||||
} |
|
||||
} |
|
||||
|
|
||||
// validateStruct receives struct type
|
|
||||
func (v *defaultValidator) validateStruct(obj interface{}) error { |
|
||||
v.lazyinit() |
|
||||
return v.validate.Struct(obj) |
|
||||
} |
|
||||
|
|
||||
// Engine returns the underlying validator engine which powers the default
|
|
||||
// Validator instance. This is useful if you want to register custom validations
|
|
||||
// or struct level validations. See validator GoDoc for more info -
|
|
||||
// https://pkg.go.dev/github.com/go-playground/validator/v10
|
|
||||
func (v *defaultValidator) Engine() interface{} { |
|
||||
v.lazyinit() |
|
||||
return v.validate |
|
||||
} |
|
||||
|
|
||||
func (v *defaultValidator) lazyinit() { |
|
||||
v.once.Do(func() { |
|
||||
v.validate = validator.New() |
|
||||
v.validate.SetTagName("binding") |
|
||||
}) |
|
||||
} |
|
||||
@ -1,62 +0,0 @@ |
|||||
// Copyright 2014 Manu Martinez-Almeida. All rights reserved.
|
|
||||
// Use of this source code is governed by a MIT style
|
|
||||
// license that can be found in the LICENSE file.
|
|
||||
|
|
||||
package binding |
|
||||
|
|
||||
import ( |
|
||||
"errors" |
|
||||
"net/http" |
|
||||
) |
|
||||
|
|
||||
const defaultMemory = 32 << 20 |
|
||||
|
|
||||
type formBinding struct{} |
|
||||
type formPostBinding struct{} |
|
||||
type formMultipartBinding struct{} |
|
||||
|
|
||||
func (formBinding) Name() string { |
|
||||
return "form" |
|
||||
} |
|
||||
|
|
||||
func (formBinding) Bind(req *http.Request, obj interface{}) error { |
|
||||
if err := req.ParseForm(); err != nil { |
|
||||
return err |
|
||||
} |
|
||||
if err := req.ParseMultipartForm(defaultMemory); err != nil && !errors.Is(err, http.ErrNotMultipart) { |
|
||||
return err |
|
||||
} |
|
||||
if err := mapForm(obj, req.Form); err != nil { |
|
||||
return err |
|
||||
} |
|
||||
return validate(obj) |
|
||||
} |
|
||||
|
|
||||
func (formPostBinding) Name() string { |
|
||||
return "form-urlencoded" |
|
||||
} |
|
||||
|
|
||||
func (formPostBinding) Bind(req *http.Request, obj interface{}) error { |
|
||||
if err := req.ParseForm(); err != nil { |
|
||||
return err |
|
||||
} |
|
||||
if err := mapForm(obj, req.PostForm); err != nil { |
|
||||
return err |
|
||||
} |
|
||||
return validate(obj) |
|
||||
} |
|
||||
|
|
||||
func (formMultipartBinding) Name() string { |
|
||||
return "multipart/form-data" |
|
||||
} |
|
||||
|
|
||||
func (formMultipartBinding) Bind(req *http.Request, obj interface{}) error { |
|
||||
if err := req.ParseMultipartForm(defaultMemory); err != nil { |
|
||||
return err |
|
||||
} |
|
||||
if err := mappingByPtr(obj, (*multipartRequest)(req), "form"); err != nil { |
|
||||
return err |
|
||||
} |
|
||||
|
|
||||
return validate(obj) |
|
||||
} |
|
||||
@ -1,404 +0,0 @@ |
|||||
// Copyright 2014 Manu Martinez-Almeida. All rights reserved.
|
|
||||
// Use of this source code is governed by a MIT style
|
|
||||
// license that can be found in the LICENSE file.
|
|
||||
|
|
||||
package binding |
|
||||
|
|
||||
import ( |
|
||||
"errors" |
|
||||
"fmt" |
|
||||
"reflect" |
|
||||
"strconv" |
|
||||
"strings" |
|
||||
"time" |
|
||||
|
|
||||
json "github.com/json-iterator/go" |
|
||||
|
|
||||
"yunyan/lego/utils" |
|
||||
) |
|
||||
|
|
||||
var ( |
|
||||
errUnknownType = errors.New("unknown type") |
|
||||
|
|
||||
// ErrConvertMapStringSlice can not covert to map[string][]string
|
|
||||
ErrConvertMapStringSlice = errors.New("can not convert to map slices of strings") |
|
||||
|
|
||||
// ErrConvertToMapString can not convert to map[string]string
|
|
||||
ErrConvertToMapString = errors.New("can not convert to map of strings") |
|
||||
) |
|
||||
|
|
||||
func mapURI(ptr interface{}, m map[string][]string) error { |
|
||||
return mapFormByTag(ptr, m, "uri") |
|
||||
} |
|
||||
|
|
||||
func mapForm(ptr interface{}, form map[string][]string) error { |
|
||||
return mapFormByTag(ptr, form, "form") |
|
||||
} |
|
||||
|
|
||||
func MapFormWithTag(ptr interface{}, form map[string][]string, tag string) error { |
|
||||
return mapFormByTag(ptr, form, tag) |
|
||||
} |
|
||||
|
|
||||
var emptyField = reflect.StructField{} |
|
||||
|
|
||||
func mapFormByTag(ptr interface{}, form map[string][]string, tag string) error { |
|
||||
// Check if ptr is a map
|
|
||||
ptrVal := reflect.ValueOf(ptr) |
|
||||
var pointed interface{} |
|
||||
if ptrVal.Kind() == reflect.Ptr { |
|
||||
ptrVal = ptrVal.Elem() |
|
||||
pointed = ptrVal.Interface() |
|
||||
} |
|
||||
if ptrVal.Kind() == reflect.Map && |
|
||||
ptrVal.Type().Key().Kind() == reflect.String { |
|
||||
if pointed != nil { |
|
||||
ptr = pointed |
|
||||
} |
|
||||
return setFormMap(ptr, form) |
|
||||
} |
|
||||
|
|
||||
return mappingByPtr(ptr, formSource(form), tag) |
|
||||
} |
|
||||
|
|
||||
// setter tries to set value on a walking by fields of a struct
|
|
||||
type setter interface { |
|
||||
TrySet(value reflect.Value, field reflect.StructField, key string, opt setOptions) (isSet bool, err error) |
|
||||
} |
|
||||
|
|
||||
type formSource map[string][]string |
|
||||
|
|
||||
var _ setter = formSource(nil) |
|
||||
|
|
||||
// TrySet tries to set a value by request's form source (like map[string][]string)
|
|
||||
func (form formSource) TrySet(value reflect.Value, field reflect.StructField, tagValue string, opt setOptions) (isSet bool, err error) { |
|
||||
return setByForm(value, field, form, tagValue, opt) |
|
||||
} |
|
||||
|
|
||||
func mappingByPtr(ptr interface{}, setter setter, tag string) error { |
|
||||
_, err := mapping(reflect.ValueOf(ptr), emptyField, setter, tag) |
|
||||
return err |
|
||||
} |
|
||||
|
|
||||
func mapping(value reflect.Value, field reflect.StructField, setter setter, tag string) (bool, error) { |
|
||||
if field.Tag.Get(tag) == "-" { // just ignoring this field
|
|
||||
return false, nil |
|
||||
} |
|
||||
|
|
||||
vKind := value.Kind() |
|
||||
|
|
||||
if vKind == reflect.Ptr { |
|
||||
var isNew bool |
|
||||
vPtr := value |
|
||||
if value.IsNil() { |
|
||||
isNew = true |
|
||||
vPtr = reflect.New(value.Type().Elem()) |
|
||||
} |
|
||||
isSet, err := mapping(vPtr.Elem(), field, setter, tag) |
|
||||
if err != nil { |
|
||||
return false, err |
|
||||
} |
|
||||
if isNew && isSet { |
|
||||
value.Set(vPtr) |
|
||||
} |
|
||||
return isSet, nil |
|
||||
} |
|
||||
|
|
||||
if vKind != reflect.Struct || !field.Anonymous { |
|
||||
ok, err := tryToSetValue(value, field, setter, tag) |
|
||||
if err != nil { |
|
||||
return false, err |
|
||||
} |
|
||||
if ok { |
|
||||
return true, nil |
|
||||
} |
|
||||
} |
|
||||
|
|
||||
if vKind == reflect.Struct { |
|
||||
tValue := value.Type() |
|
||||
|
|
||||
var isSet bool |
|
||||
for i := 0; i < value.NumField(); i++ { |
|
||||
sf := tValue.Field(i) |
|
||||
if sf.PkgPath != "" && !sf.Anonymous { // unexported
|
|
||||
continue |
|
||||
} |
|
||||
ok, err := mapping(value.Field(i), sf, setter, tag) |
|
||||
if err != nil { |
|
||||
return false, err |
|
||||
} |
|
||||
isSet = isSet || ok |
|
||||
} |
|
||||
return isSet, nil |
|
||||
} |
|
||||
return false, nil |
|
||||
} |
|
||||
|
|
||||
type setOptions struct { |
|
||||
isDefaultExists bool |
|
||||
defaultValue string |
|
||||
} |
|
||||
|
|
||||
func tryToSetValue(value reflect.Value, field reflect.StructField, setter setter, tag string) (bool, error) { |
|
||||
var tagValue string |
|
||||
var setOpt setOptions |
|
||||
|
|
||||
tagValue = field.Tag.Get(tag) |
|
||||
tagValue, opts := head(tagValue, ",") |
|
||||
|
|
||||
if tagValue == "" { // default value is FieldName
|
|
||||
tagValue = field.Name |
|
||||
} |
|
||||
if tagValue == "" { // when field is "emptyField" variable
|
|
||||
return false, nil |
|
||||
} |
|
||||
|
|
||||
var opt string |
|
||||
for len(opts) > 0 { |
|
||||
opt, opts = head(opts, ",") |
|
||||
|
|
||||
if k, v := head(opt, "="); k == "default" { |
|
||||
setOpt.isDefaultExists = true |
|
||||
setOpt.defaultValue = v |
|
||||
} |
|
||||
} |
|
||||
|
|
||||
return setter.TrySet(value, field, tagValue, setOpt) |
|
||||
} |
|
||||
|
|
||||
func setByForm(value reflect.Value, field reflect.StructField, form map[string][]string, tagValue string, opt setOptions) (isSet bool, err error) { |
|
||||
vs, ok := form[tagValue] |
|
||||
if !ok && !opt.isDefaultExists { |
|
||||
return false, nil |
|
||||
} |
|
||||
|
|
||||
switch value.Kind() { |
|
||||
case reflect.Slice: |
|
||||
if !ok { |
|
||||
vs = []string{opt.defaultValue} |
|
||||
} |
|
||||
return true, setSlice(vs, value, field) |
|
||||
case reflect.Array: |
|
||||
if !ok { |
|
||||
vs = []string{opt.defaultValue} |
|
||||
} |
|
||||
if len(vs) != value.Len() { |
|
||||
return false, fmt.Errorf("%q is not valid value for %s", vs, value.Type().String()) |
|
||||
} |
|
||||
return true, setArray(vs, value, field) |
|
||||
default: |
|
||||
var val string |
|
||||
if !ok { |
|
||||
val = opt.defaultValue |
|
||||
} |
|
||||
|
|
||||
if len(vs) > 0 { |
|
||||
val = vs[0] |
|
||||
} |
|
||||
return true, setWithProperType(val, value, field) |
|
||||
} |
|
||||
} |
|
||||
|
|
||||
func setWithProperType(val string, value reflect.Value, field reflect.StructField) error { |
|
||||
switch value.Kind() { |
|
||||
case reflect.Int: |
|
||||
return setIntField(val, 0, value) |
|
||||
case reflect.Int8: |
|
||||
return setIntField(val, 8, value) |
|
||||
case reflect.Int16: |
|
||||
return setIntField(val, 16, value) |
|
||||
case reflect.Int32: |
|
||||
return setIntField(val, 32, value) |
|
||||
case reflect.Int64: |
|
||||
switch value.Interface().(type) { |
|
||||
case time.Duration: |
|
||||
return setTimeDuration(val, value) |
|
||||
} |
|
||||
return setIntField(val, 64, value) |
|
||||
case reflect.Uint: |
|
||||
return setUintField(val, 0, value) |
|
||||
case reflect.Uint8: |
|
||||
return setUintField(val, 8, value) |
|
||||
case reflect.Uint16: |
|
||||
return setUintField(val, 16, value) |
|
||||
case reflect.Uint32: |
|
||||
return setUintField(val, 32, value) |
|
||||
case reflect.Uint64: |
|
||||
return setUintField(val, 64, value) |
|
||||
case reflect.Bool: |
|
||||
return setBoolField(val, value) |
|
||||
case reflect.Float32: |
|
||||
return setFloatField(val, 32, value) |
|
||||
case reflect.Float64: |
|
||||
return setFloatField(val, 64, value) |
|
||||
case reflect.String: |
|
||||
value.SetString(val) |
|
||||
case reflect.Struct: |
|
||||
switch value.Interface().(type) { |
|
||||
case time.Time: |
|
||||
return setTimeField(val, field, value) |
|
||||
} |
|
||||
return json.Unmarshal(utils.StringToBytes(val), value.Addr().Interface()) |
|
||||
case reflect.Map: |
|
||||
return json.Unmarshal(utils.StringToBytes(val), value.Addr().Interface()) |
|
||||
default: |
|
||||
return errUnknownType |
|
||||
} |
|
||||
return nil |
|
||||
} |
|
||||
|
|
||||
func setIntField(val string, bitSize int, field reflect.Value) error { |
|
||||
if val == "" { |
|
||||
val = "0" |
|
||||
} |
|
||||
intVal, err := strconv.ParseInt(val, 10, bitSize) |
|
||||
if err == nil { |
|
||||
field.SetInt(intVal) |
|
||||
} |
|
||||
return err |
|
||||
} |
|
||||
|
|
||||
func setUintField(val string, bitSize int, field reflect.Value) error { |
|
||||
if val == "" { |
|
||||
val = "0" |
|
||||
} |
|
||||
uintVal, err := strconv.ParseUint(val, 10, bitSize) |
|
||||
if err == nil { |
|
||||
field.SetUint(uintVal) |
|
||||
} |
|
||||
return err |
|
||||
} |
|
||||
|
|
||||
func setBoolField(val string, field reflect.Value) error { |
|
||||
if val == "" { |
|
||||
val = "false" |
|
||||
} |
|
||||
boolVal, err := strconv.ParseBool(val) |
|
||||
if err == nil { |
|
||||
field.SetBool(boolVal) |
|
||||
} |
|
||||
return err |
|
||||
} |
|
||||
|
|
||||
func setFloatField(val string, bitSize int, field reflect.Value) error { |
|
||||
if val == "" { |
|
||||
val = "0.0" |
|
||||
} |
|
||||
floatVal, err := strconv.ParseFloat(val, bitSize) |
|
||||
if err == nil { |
|
||||
field.SetFloat(floatVal) |
|
||||
} |
|
||||
return err |
|
||||
} |
|
||||
|
|
||||
func setTimeField(val string, structField reflect.StructField, value reflect.Value) error { |
|
||||
timeFormat := structField.Tag.Get("time_format") |
|
||||
if timeFormat == "" { |
|
||||
timeFormat = time.RFC3339 |
|
||||
} |
|
||||
|
|
||||
switch tf := strings.ToLower(timeFormat); tf { |
|
||||
case "unix", "unixnano": |
|
||||
tv, err := strconv.ParseInt(val, 10, 64) |
|
||||
if err != nil { |
|
||||
return err |
|
||||
} |
|
||||
|
|
||||
d := time.Duration(1) |
|
||||
if tf == "unixnano" { |
|
||||
d = time.Second |
|
||||
} |
|
||||
|
|
||||
t := time.Unix(tv/int64(d), tv%int64(d)) |
|
||||
value.Set(reflect.ValueOf(t)) |
|
||||
return nil |
|
||||
} |
|
||||
|
|
||||
if val == "" { |
|
||||
value.Set(reflect.ValueOf(time.Time{})) |
|
||||
return nil |
|
||||
} |
|
||||
|
|
||||
l := time.Local |
|
||||
if isUTC, _ := strconv.ParseBool(structField.Tag.Get("time_utc")); isUTC { |
|
||||
l = time.UTC |
|
||||
} |
|
||||
|
|
||||
if locTag := structField.Tag.Get("time_location"); locTag != "" { |
|
||||
loc, err := time.LoadLocation(locTag) |
|
||||
if err != nil { |
|
||||
return err |
|
||||
} |
|
||||
l = loc |
|
||||
} |
|
||||
|
|
||||
t, err := time.ParseInLocation(timeFormat, val, l) |
|
||||
if err != nil { |
|
||||
return err |
|
||||
} |
|
||||
|
|
||||
value.Set(reflect.ValueOf(t)) |
|
||||
return nil |
|
||||
} |
|
||||
|
|
||||
func setArray(vals []string, value reflect.Value, field reflect.StructField) error { |
|
||||
for i, s := range vals { |
|
||||
err := setWithProperType(s, value.Index(i), field) |
|
||||
if err != nil { |
|
||||
return err |
|
||||
} |
|
||||
} |
|
||||
return nil |
|
||||
} |
|
||||
|
|
||||
func setSlice(vals []string, value reflect.Value, field reflect.StructField) error { |
|
||||
slice := reflect.MakeSlice(value.Type(), len(vals), len(vals)) |
|
||||
err := setArray(vals, slice, field) |
|
||||
if err != nil { |
|
||||
return err |
|
||||
} |
|
||||
value.Set(slice) |
|
||||
return nil |
|
||||
} |
|
||||
|
|
||||
func setTimeDuration(val string, value reflect.Value) error { |
|
||||
d, err := time.ParseDuration(val) |
|
||||
if err != nil { |
|
||||
return err |
|
||||
} |
|
||||
value.Set(reflect.ValueOf(d)) |
|
||||
return nil |
|
||||
} |
|
||||
|
|
||||
func head(str, sep string) (head string, tail string) { |
|
||||
idx := strings.Index(str, sep) |
|
||||
if idx < 0 { |
|
||||
return str, "" |
|
||||
} |
|
||||
return str[:idx], str[idx+len(sep):] |
|
||||
} |
|
||||
|
|
||||
func setFormMap(ptr interface{}, form map[string][]string) error { |
|
||||
el := reflect.TypeOf(ptr).Elem() |
|
||||
|
|
||||
if el.Kind() == reflect.Slice { |
|
||||
ptrMap, ok := ptr.(map[string][]string) |
|
||||
if !ok { |
|
||||
return ErrConvertMapStringSlice |
|
||||
} |
|
||||
for k, v := range form { |
|
||||
ptrMap[k] = v |
|
||||
} |
|
||||
|
|
||||
return nil |
|
||||
} |
|
||||
|
|
||||
ptrMap, ok := ptr.(map[string]string) |
|
||||
if !ok { |
|
||||
return ErrConvertToMapString |
|
||||
} |
|
||||
for k, v := range form { |
|
||||
ptrMap[k] = v[len(v)-1] // pick last
|
|
||||
} |
|
||||
|
|
||||
return nil |
|
||||
} |
|
||||
@ -1,34 +0,0 @@ |
|||||
package binding |
|
||||
|
|
||||
import ( |
|
||||
"net/http" |
|
||||
"net/textproto" |
|
||||
"reflect" |
|
||||
) |
|
||||
|
|
||||
type headerBinding struct{} |
|
||||
|
|
||||
func (headerBinding) Name() string { |
|
||||
return "header" |
|
||||
} |
|
||||
|
|
||||
func (headerBinding) Bind(req *http.Request, obj interface{}) error { |
|
||||
|
|
||||
if err := mapHeader(obj, req.Header); err != nil { |
|
||||
return err |
|
||||
} |
|
||||
|
|
||||
return validate(obj) |
|
||||
} |
|
||||
|
|
||||
func mapHeader(ptr interface{}, h map[string][]string) error { |
|
||||
return mappingByPtr(ptr, headerSource(h), "header") |
|
||||
} |
|
||||
|
|
||||
type headerSource map[string][]string |
|
||||
|
|
||||
var _ setter = headerSource(nil) |
|
||||
|
|
||||
func (hs headerSource) TrySet(value reflect.Value, field reflect.StructField, tagValue string, opt setOptions) (bool, error) { |
|
||||
return setByForm(value, field, hs, textproto.CanonicalMIMEHeaderKey(tagValue), opt) |
|
||||
} |
|
||||
@ -1,56 +0,0 @@ |
|||||
// Copyright 2014 Manu Martinez-Almeida. All rights reserved.
|
|
||||
// Use of this source code is governed by a MIT style
|
|
||||
// license that can be found in the LICENSE file.
|
|
||||
|
|
||||
package binding |
|
||||
|
|
||||
import ( |
|
||||
"bytes" |
|
||||
"errors" |
|
||||
"io" |
|
||||
"net/http" |
|
||||
|
|
||||
json "github.com/json-iterator/go" |
|
||||
) |
|
||||
|
|
||||
// EnableDecoderUseNumber is used to call the UseNumber method on the JSON
|
|
||||
// Decoder instance. UseNumber causes the Decoder to unmarshal a number into an
|
|
||||
// interface{} as a Number instead of as a float64.
|
|
||||
var EnableDecoderUseNumber = false |
|
||||
|
|
||||
// EnableDecoderDisallowUnknownFields is used to call the DisallowUnknownFields method
|
|
||||
// on the JSON Decoder instance. DisallowUnknownFields causes the Decoder to
|
|
||||
// return an error when the destination is a struct and the input contains object
|
|
||||
// keys which do not match any non-ignored, exported fields in the destination.
|
|
||||
var EnableDecoderDisallowUnknownFields = false |
|
||||
|
|
||||
type jsonBinding struct{} |
|
||||
|
|
||||
func (jsonBinding) Name() string { |
|
||||
return "json" |
|
||||
} |
|
||||
|
|
||||
func (jsonBinding) Bind(req *http.Request, obj interface{}) error { |
|
||||
if req == nil || req.Body == nil { |
|
||||
return errors.New("invalid request") |
|
||||
} |
|
||||
return decodeJSON(req.Body, obj) |
|
||||
} |
|
||||
|
|
||||
func (jsonBinding) BindBody(body []byte, obj interface{}) error { |
|
||||
return decodeJSON(bytes.NewReader(body), obj) |
|
||||
} |
|
||||
|
|
||||
func decodeJSON(r io.Reader, obj interface{}) error { |
|
||||
decoder := json.NewDecoder(r) |
|
||||
if EnableDecoderUseNumber { |
|
||||
decoder.UseNumber() |
|
||||
} |
|
||||
if EnableDecoderDisallowUnknownFields { |
|
||||
decoder.DisallowUnknownFields() |
|
||||
} |
|
||||
if err := decoder.Decode(obj); err != nil { |
|
||||
return err |
|
||||
} |
|
||||
return validate(obj) |
|
||||
} |
|
||||
@ -1,31 +0,0 @@ |
|||||
package binding |
|
||||
|
|
||||
import ( |
|
||||
"bytes" |
|
||||
"io" |
|
||||
"net/http" |
|
||||
|
|
||||
"github.com/ugorji/go/codec" |
|
||||
) |
|
||||
|
|
||||
type msgpackBinding struct{} |
|
||||
|
|
||||
func (msgpackBinding) Name() string { |
|
||||
return "msgpack" |
|
||||
} |
|
||||
|
|
||||
func (msgpackBinding) Bind(req *http.Request, obj interface{}) error { |
|
||||
return decodeMsgPack(req.Body, obj) |
|
||||
} |
|
||||
|
|
||||
func (msgpackBinding) BindBody(body []byte, obj interface{}) error { |
|
||||
return decodeMsgPack(bytes.NewReader(body), obj) |
|
||||
} |
|
||||
|
|
||||
func decodeMsgPack(r io.Reader, obj interface{}) error { |
|
||||
cdc := new(codec.MsgpackHandle) |
|
||||
if err := codec.NewDecoder(r, cdc).Decode(&obj); err != nil { |
|
||||
return err |
|
||||
} |
|
||||
return validate(obj) |
|
||||
} |
|
||||
@ -1,74 +0,0 @@ |
|||||
// Copyright 2019 Gin Core Team. All rights reserved.
|
|
||||
// Use of this source code is governed by a MIT style
|
|
||||
// license that can be found in the LICENSE file.
|
|
||||
|
|
||||
package binding |
|
||||
|
|
||||
import ( |
|
||||
"errors" |
|
||||
"mime/multipart" |
|
||||
"net/http" |
|
||||
"reflect" |
|
||||
) |
|
||||
|
|
||||
type multipartRequest http.Request |
|
||||
|
|
||||
var _ setter = (*multipartRequest)(nil) |
|
||||
|
|
||||
var ( |
|
||||
// ErrMultiFileHeader multipart.FileHeader invalid
|
|
||||
ErrMultiFileHeader = errors.New("unsupported field type for multipart.FileHeader") |
|
||||
|
|
||||
// ErrMultiFileHeaderLenInvalid array for []*multipart.FileHeader len invalid
|
|
||||
ErrMultiFileHeaderLenInvalid = errors.New("unsupported len of array for []*multipart.FileHeader") |
|
||||
) |
|
||||
|
|
||||
// TrySet tries to set a value by the multipart request with the binding a form file
|
|
||||
func (r *multipartRequest) TrySet(value reflect.Value, field reflect.StructField, key string, opt setOptions) (bool, error) { |
|
||||
if files := r.MultipartForm.File[key]; len(files) != 0 { |
|
||||
return setByMultipartFormFile(value, field, files) |
|
||||
} |
|
||||
|
|
||||
return setByForm(value, field, r.MultipartForm.Value, key, opt) |
|
||||
} |
|
||||
|
|
||||
func setByMultipartFormFile(value reflect.Value, field reflect.StructField, files []*multipart.FileHeader) (isSet bool, err error) { |
|
||||
switch value.Kind() { |
|
||||
case reflect.Ptr: |
|
||||
switch value.Interface().(type) { |
|
||||
case *multipart.FileHeader: |
|
||||
value.Set(reflect.ValueOf(files[0])) |
|
||||
return true, nil |
|
||||
} |
|
||||
case reflect.Struct: |
|
||||
switch value.Interface().(type) { |
|
||||
case multipart.FileHeader: |
|
||||
value.Set(reflect.ValueOf(*files[0])) |
|
||||
return true, nil |
|
||||
} |
|
||||
case reflect.Slice: |
|
||||
slice := reflect.MakeSlice(value.Type(), len(files), len(files)) |
|
||||
isSet, err = setArrayOfMultipartFormFiles(slice, field, files) |
|
||||
if err != nil || !isSet { |
|
||||
return isSet, err |
|
||||
} |
|
||||
value.Set(slice) |
|
||||
return true, nil |
|
||||
case reflect.Array: |
|
||||
return setArrayOfMultipartFormFiles(value, field, files) |
|
||||
} |
|
||||
return false, ErrMultiFileHeader |
|
||||
} |
|
||||
|
|
||||
func setArrayOfMultipartFormFiles(value reflect.Value, field reflect.StructField, files []*multipart.FileHeader) (isSet bool, err error) { |
|
||||
if value.Len() != len(files) { |
|
||||
return false, ErrMultiFileHeaderLenInvalid |
|
||||
} |
|
||||
for i := range files { |
|
||||
set, err := setByMultipartFormFile(value.Index(i), field, files[i:i+1]) |
|
||||
if err != nil || !set { |
|
||||
return set, err |
|
||||
} |
|
||||
} |
|
||||
return true, nil |
|
||||
} |
|
||||
@ -1,35 +0,0 @@ |
|||||
package binding |
|
||||
|
|
||||
import ( |
|
||||
"errors" |
|
||||
"io/ioutil" |
|
||||
"net/http" |
|
||||
|
|
||||
"google.golang.org/protobuf/proto" |
|
||||
) |
|
||||
|
|
||||
type protobufBinding struct{} |
|
||||
|
|
||||
func (protobufBinding) Name() string { |
|
||||
return "protobuf" |
|
||||
} |
|
||||
|
|
||||
func (b protobufBinding) Bind(req *http.Request, obj interface{}) error { |
|
||||
buf, err := ioutil.ReadAll(req.Body) |
|
||||
if err != nil { |
|
||||
return err |
|
||||
} |
|
||||
return b.BindBody(buf, obj) |
|
||||
} |
|
||||
|
|
||||
func (protobufBinding) BindBody(body []byte, obj interface{}) error { |
|
||||
msg, ok := obj.(proto.Message) |
|
||||
if !ok { |
|
||||
return errors.New("obj is not ProtoMessage") |
|
||||
} |
|
||||
if err := proto.Unmarshal(body, msg); err != nil { |
|
||||
return err |
|
||||
} |
|
||||
return nil |
|
||||
|
|
||||
} |
|
||||
@ -1,17 +0,0 @@ |
|||||
package binding |
|
||||
|
|
||||
import "net/http" |
|
||||
|
|
||||
type queryBinding struct{} |
|
||||
|
|
||||
func (queryBinding) Name() string { |
|
||||
return "query" |
|
||||
} |
|
||||
|
|
||||
func (queryBinding) Bind(req *http.Request, obj interface{}) error { |
|
||||
values := req.URL.Query() |
|
||||
if err := mapForm(obj, values); err != nil { |
|
||||
return err |
|
||||
} |
|
||||
return validate(obj) |
|
||||
} |
|
||||
@ -1,14 +0,0 @@ |
|||||
package binding |
|
||||
|
|
||||
type uriBinding struct{} |
|
||||
|
|
||||
func (uriBinding) Name() string { |
|
||||
return "uri" |
|
||||
} |
|
||||
|
|
||||
func (uriBinding) BindUri(m map[string][]string, obj interface{}) error { |
|
||||
if err := mapURI(obj, m); err != nil { |
|
||||
return err |
|
||||
} |
|
||||
return validate(obj) |
|
||||
} |
|
||||
@ -1,33 +0,0 @@ |
|||||
// Copyright 2014 Manu Martinez-Almeida. All rights reserved.
|
|
||||
// Use of this source code is governed by a MIT style
|
|
||||
// license that can be found in the LICENSE file.
|
|
||||
|
|
||||
package binding |
|
||||
|
|
||||
import ( |
|
||||
"bytes" |
|
||||
"encoding/xml" |
|
||||
"io" |
|
||||
"net/http" |
|
||||
) |
|
||||
|
|
||||
type xmlBinding struct{} |
|
||||
|
|
||||
func (xmlBinding) Name() string { |
|
||||
return "xml" |
|
||||
} |
|
||||
|
|
||||
func (xmlBinding) Bind(req *http.Request, obj interface{}) error { |
|
||||
return decodeXML(req.Body, obj) |
|
||||
} |
|
||||
|
|
||||
func (xmlBinding) BindBody(body []byte, obj interface{}) error { |
|
||||
return decodeXML(bytes.NewReader(body), obj) |
|
||||
} |
|
||||
func decodeXML(r io.Reader, obj interface{}) error { |
|
||||
decoder := xml.NewDecoder(r) |
|
||||
if err := decoder.Decode(obj); err != nil { |
|
||||
return err |
|
||||
} |
|
||||
return validate(obj) |
|
||||
} |
|
||||
@ -1,31 +0,0 @@ |
|||||
package binding |
|
||||
|
|
||||
import ( |
|
||||
"bytes" |
|
||||
"io" |
|
||||
"net/http" |
|
||||
|
|
||||
"gopkg.in/yaml.v2" |
|
||||
) |
|
||||
|
|
||||
type yamlBinding struct{} |
|
||||
|
|
||||
func (yamlBinding) Name() string { |
|
||||
return "yaml" |
|
||||
} |
|
||||
|
|
||||
func (yamlBinding) Bind(req *http.Request, obj interface{}) error { |
|
||||
return decodeYAML(req.Body, obj) |
|
||||
} |
|
||||
|
|
||||
func (yamlBinding) BindBody(body []byte, obj interface{}) error { |
|
||||
return decodeYAML(bytes.NewReader(body), obj) |
|
||||
} |
|
||||
|
|
||||
func decodeYAML(r io.Reader, obj interface{}) error { |
|
||||
decoder := yaml.NewDecoder(r) |
|
||||
if err := decoder.Decode(obj); err != nil { |
|
||||
return err |
|
||||
} |
|
||||
return validate(obj) |
|
||||
} |
|
||||
@ -1,153 +0,0 @@ |
|||||
package gin |
|
||||
|
|
||||
import ( |
|
||||
"fmt" |
|
||||
"net/http" |
|
||||
"reflect" |
|
||||
"sort" |
|
||||
"strings" |
|
||||
|
|
||||
"yunyan/lego/sys/gin/engine" |
|
||||
"yunyan/lego/utils/crypto/md5" |
|
||||
) |
|
||||
|
|
||||
/* |
|
||||
系统描述:开源gin框架的重构版本 |
|
||||
*/ |
|
||||
type ISys interface { |
|
||||
engine.IRoutes |
|
||||
HandleContext(c *engine.Context) |
|
||||
LoadHTMLGlob(pattern string) |
|
||||
Close() (err error) |
|
||||
} |
|
||||
|
|
||||
var defsys ISys |
|
||||
|
|
||||
func OnInit(config map[string]interface{}, opt ...Option) (err error) { |
|
||||
var option *Options |
|
||||
if option, err = newOptions(config, opt...); err != nil { |
|
||||
return |
|
||||
} |
|
||||
defsys, err = newSys(option) |
|
||||
return |
|
||||
} |
|
||||
|
|
||||
func NewSys(opt ...Option) (sys ISys, err error) { |
|
||||
var option *Options |
|
||||
if option, err = newOptionsByOption(opt...); err != nil { |
|
||||
return |
|
||||
} |
|
||||
sys, err = newSys(option) |
|
||||
return |
|
||||
} |
|
||||
|
|
||||
func LoadHTMLGlob(pattern string) { |
|
||||
defsys.LoadHTMLGlob(pattern) |
|
||||
} |
|
||||
|
|
||||
func HandleContext(c *engine.Context) { |
|
||||
defsys.HandleContext(c) |
|
||||
} |
|
||||
|
|
||||
func Close() (err error) { |
|
||||
return defsys.Close() |
|
||||
} |
|
||||
func NoRoute(handlers ...engine.HandlerFunc) { |
|
||||
defsys.NoRoute(handlers...) |
|
||||
} |
|
||||
func Use(handlers ...engine.HandlerFunc) engine.IRoutes { |
|
||||
return defsys.Use(handlers...) |
|
||||
} |
|
||||
func Handle(httpMethod string, relativePath string, handlers ...engine.HandlerFunc) engine.IRoutes { |
|
||||
return defsys.Handle(httpMethod, relativePath, handlers...) |
|
||||
} |
|
||||
func Any(relativePath string, handlers ...engine.HandlerFunc) engine.IRoutes { |
|
||||
return defsys.Any(relativePath, handlers...) |
|
||||
} |
|
||||
func GET(httpMethod string, handlers ...engine.HandlerFunc) engine.IRoutes { |
|
||||
return defsys.GET(httpMethod, handlers...) |
|
||||
} |
|
||||
func POST(httpMethod string, handlers ...engine.HandlerFunc) engine.IRoutes { |
|
||||
return defsys.POST(httpMethod, handlers...) |
|
||||
} |
|
||||
func DELETE(httpMethod string, handlers ...engine.HandlerFunc) engine.IRoutes { |
|
||||
return defsys.DELETE(httpMethod, handlers...) |
|
||||
} |
|
||||
func PATCH(httpMethod string, handlers ...engine.HandlerFunc) engine.IRoutes { |
|
||||
return defsys.PATCH(httpMethod, handlers...) |
|
||||
} |
|
||||
func PUT(httpMethod string, handlers ...engine.HandlerFunc) engine.IRoutes { |
|
||||
return defsys.PUT(httpMethod, handlers...) |
|
||||
} |
|
||||
func OPTIONS(httpMethod string, handlers ...engine.HandlerFunc) engine.IRoutes { |
|
||||
return defsys.OPTIONS(httpMethod, handlers...) |
|
||||
} |
|
||||
func HEAD(httpMethod string, handlers ...engine.HandlerFunc) engine.IRoutes { |
|
||||
return defsys.HEAD(httpMethod, handlers...) |
|
||||
} |
|
||||
func StaticFile(relativePath string, filepath string) engine.IRoutes { |
|
||||
return defsys.StaticFile(relativePath, filepath) |
|
||||
} |
|
||||
func StaticFileFS(relativePath string, filepath string, fs http.FileSystem) engine.IRoutes { |
|
||||
return defsys.StaticFileFS(relativePath, filepath, fs) |
|
||||
} |
|
||||
|
|
||||
func Static(relativePath string, root string) engine.IRoutes { |
|
||||
return defsys.Static(relativePath, root) |
|
||||
} |
|
||||
|
|
||||
func StaticFS(relativePath string, fs http.FileSystem) engine.IRoutes { |
|
||||
return defsys.StaticFS(relativePath, fs) |
|
||||
} |
|
||||
|
|
||||
// 签名接口
|
|
||||
func ParamSign(key string, param map[string]interface{}) (origin, sign string) { |
|
||||
var keys []string |
|
||||
for k, _ := range param { |
|
||||
keys = append(keys, k) |
|
||||
} |
|
||||
sort.Strings(keys) |
|
||||
builder := strings.Builder{} |
|
||||
for _, v := range keys { |
|
||||
builder.WriteString(v) |
|
||||
builder.WriteString("=") |
|
||||
switch reflect.TypeOf(param[v]).Kind() { |
|
||||
case reflect.Int, |
|
||||
reflect.Int8, |
|
||||
reflect.Int16, |
|
||||
reflect.Int32, |
|
||||
reflect.Int64, |
|
||||
reflect.Uint, |
|
||||
reflect.Uint8, |
|
||||
reflect.Uint16, |
|
||||
reflect.Uint32, |
|
||||
reflect.Uint64: |
|
||||
builder.WriteString(fmt.Sprintf("%d", param[v])) |
|
||||
break |
|
||||
case reflect.Float32, |
|
||||
reflect.Float64: |
|
||||
builder.WriteString(fmt.Sprintf("%v", param[v])) |
|
||||
case reflect.Bool: |
|
||||
builder.WriteString(fmt.Sprintf("%v", param[v])) |
|
||||
case reflect.Slice, reflect.Array: |
|
||||
s := reflect.ValueOf(param[v]) |
|
||||
valueStr := "" |
|
||||
for i := 0; i < s.Len(); i++ { |
|
||||
valueStr += fmt.Sprintf("%v,", s.Index(i).Interface()) |
|
||||
} |
|
||||
if s.Len() > 0 { |
|
||||
valueStr = valueStr[0 : len(valueStr)-1] |
|
||||
} |
|
||||
builder.WriteString(fmt.Sprintf("%s", valueStr)) |
|
||||
break |
|
||||
default: |
|
||||
builder.WriteString(fmt.Sprintf("%s", param[v])) |
|
||||
break |
|
||||
} |
|
||||
builder.WriteString("&") |
|
||||
} |
|
||||
builder.WriteString("key=" + key) |
|
||||
origin = builder.String() |
|
||||
sign = md5.MD5EncToLower(origin) |
|
||||
return |
|
||||
} |
|
||||
@ -1,883 +0,0 @@ |
|||||
package engine |
|
||||
|
|
||||
import ( |
|
||||
"errors" |
|
||||
"io" |
|
||||
"io/ioutil" |
|
||||
"math" |
|
||||
"mime/multipart" |
|
||||
"net" |
|
||||
"net/http" |
|
||||
"net/url" |
|
||||
"os" |
|
||||
"strings" |
|
||||
"sync" |
|
||||
"time" |
|
||||
|
|
||||
"yunyan/lego/sys/gin/binding" |
|
||||
"yunyan/lego/sys/gin/render" |
|
||||
"yunyan/lego/sys/log" |
|
||||
) |
|
||||
|
|
||||
const ( |
|
||||
MIMEJSON = binding.MIMEJSON |
|
||||
MIMEHTML = binding.MIMEHTML |
|
||||
MIMEXML = binding.MIMEXML |
|
||||
MIMEXML2 = binding.MIMEXML2 |
|
||||
MIMEPlain = binding.MIMEPlain |
|
||||
MIMEPOSTForm = binding.MIMEPOSTForm |
|
||||
MIMEMultipartPOSTForm = binding.MIMEMultipartPOSTForm |
|
||||
MIMEYAML = binding.MIMEYAML |
|
||||
) |
|
||||
const abortIndex int8 = math.MaxInt8 >> 1 |
|
||||
|
|
||||
func newContext(log log.ILogger, engine *Engine, params *Params, skippedNodes *[]skippedNode) *Context { |
|
||||
return &Context{ |
|
||||
Log: log, |
|
||||
engine: engine, |
|
||||
params: params, |
|
||||
skippedNodes: skippedNodes, |
|
||||
writermem: ResponseWriter{log: log}, |
|
||||
} |
|
||||
} |
|
||||
|
|
||||
type Context struct { |
|
||||
Log log.ILogger |
|
||||
engine *Engine |
|
||||
writermem ResponseWriter |
|
||||
Request *http.Request |
|
||||
Writer IResponseWriter |
|
||||
Params Params |
|
||||
handlers HandlersChain |
|
||||
index int8 |
|
||||
fullPath string |
|
||||
params *Params |
|
||||
skippedNodes *[]skippedNode |
|
||||
mu sync.RWMutex |
|
||||
Keys map[string]interface{} |
|
||||
Errors errorMsgs |
|
||||
Accepted []string |
|
||||
queryCache url.Values |
|
||||
formCache url.Values |
|
||||
sameSite http.SameSite |
|
||||
} |
|
||||
|
|
||||
func (this *Context) Copy() *Context { |
|
||||
cp := Context{ |
|
||||
writermem: this.writermem, |
|
||||
Request: this.Request, |
|
||||
Params: this.Params, |
|
||||
engine: this.engine, |
|
||||
} |
|
||||
cp.writermem.ResponseWriter = nil |
|
||||
cp.Writer = &cp.writermem |
|
||||
cp.index = abortIndex |
|
||||
cp.handlers = nil |
|
||||
cp.Keys = map[string]interface{}{} |
|
||||
for k, v := range this.Keys { |
|
||||
cp.Keys[k] = v |
|
||||
} |
|
||||
paramCopy := make([]Param, len(cp.Params)) |
|
||||
copy(paramCopy, cp.Params) |
|
||||
cp.Params = paramCopy |
|
||||
return &cp |
|
||||
} |
|
||||
|
|
||||
func (this *Context) HandlerName() string { |
|
||||
return nameOfFunction(this.handlers.Last()) |
|
||||
} |
|
||||
|
|
||||
func (this *Context) HandlerNames() []string { |
|
||||
hn := make([]string, 0, len(this.handlers)) |
|
||||
for _, val := range this.handlers { |
|
||||
hn = append(hn, nameOfFunction(val)) |
|
||||
} |
|
||||
return hn |
|
||||
} |
|
||||
|
|
||||
func (this *Context) Handler() HandlerFunc { |
|
||||
return this.handlers.Last() |
|
||||
} |
|
||||
|
|
||||
/* |
|
||||
FullPath 返回匹配的路由完整路径。 对于未找到的路线 |
|
||||
返回一个空字符串。 |
|
||||
*/ |
|
||||
func (c *Context) FullPath() string { |
|
||||
return c.fullPath |
|
||||
} |
|
||||
|
|
||||
func (this *Context) Next() { |
|
||||
this.index++ |
|
||||
for this.index < int8(len(this.handlers)) { |
|
||||
this.handlers[this.index](this) |
|
||||
this.index++ |
|
||||
} |
|
||||
} |
|
||||
|
|
||||
/* |
|
||||
如果当前上下文被中止,IsAborted 返回 true。 |
|
||||
*/ |
|
||||
func (this *Context) IsAborted() bool { |
|
||||
return this.index >= abortIndex |
|
||||
} |
|
||||
|
|
||||
/* |
|
||||
Abort 防止挂起的处理程序被调用。 请注意,这不会停止当前处理程序。 |
|
||||
假设你有一个授权中间件来验证当前请求是否被授权。 |
|
||||
如果授权失败(例如:密码不匹配),调用 Abort 以确保剩余的 handlers |
|
||||
因为这个请求没有被调用。 |
|
||||
*/ |
|
||||
func (this *Context) Abort() { |
|
||||
this.index = abortIndex |
|
||||
} |
|
||||
|
|
||||
/* |
|
||||
AbortWithStatus 调用 `Abort()` 并使用指定的状态代码写入标头。 |
|
||||
例如,验证请求失败的尝试可以使用:context.AbortWithStatus(401)。 |
|
||||
*/ |
|
||||
func (this *Context) AbortWithStatus(code int) { |
|
||||
this.Status(code) |
|
||||
this.Writer.WriteHeaderNow() |
|
||||
this.Abort() |
|
||||
} |
|
||||
|
|
||||
func (this *Context) AbortWithStatusJSON(code int, jsonObj interface{}) { |
|
||||
this.Abort() |
|
||||
this.JSON(code, jsonObj) |
|
||||
} |
|
||||
|
|
||||
func (this *Context) AbortWithError(code int, err error) *Error { |
|
||||
this.AbortWithStatus(code) |
|
||||
return this.Error(err) |
|
||||
} |
|
||||
|
|
||||
func (this *Context) Set(key string, value interface{}) { |
|
||||
this.mu.Lock() |
|
||||
if this.Keys == nil { |
|
||||
this.Keys = make(map[string]interface{}) |
|
||||
} |
|
||||
|
|
||||
this.Keys[key] = value |
|
||||
this.mu.Unlock() |
|
||||
} |
|
||||
|
|
||||
func (this *Context) SetUserId(uid string) { |
|
||||
this.Set("UserId", uid) |
|
||||
} |
|
||||
|
|
||||
func (this *Context) Get(key string) (value interface{}, exists bool) { |
|
||||
this.mu.RLock() |
|
||||
value, exists = this.Keys[key] |
|
||||
this.mu.RUnlock() |
|
||||
return |
|
||||
} |
|
||||
|
|
||||
/* |
|
||||
如果存在,MustGet 返回给定键的值,否则抛出异常。 |
|
||||
*/ |
|
||||
func (this *Context) MustGet(key string) interface{} { |
|
||||
if value, exists := this.Get(key); exists { |
|
||||
return value |
|
||||
} |
|
||||
panic("Key \"" + key + "\" does not exist") |
|
||||
} |
|
||||
|
|
||||
func (this *Context) GetString(key string) (s string) { |
|
||||
if val, ok := this.Get(key); ok && val != nil { |
|
||||
s, _ = val.(string) |
|
||||
} |
|
||||
return |
|
||||
} |
|
||||
|
|
||||
func (this *Context) GetBool(key string) (b bool) { |
|
||||
if val, ok := this.Get(key); ok && val != nil { |
|
||||
b, _ = val.(bool) |
|
||||
} |
|
||||
return |
|
||||
} |
|
||||
|
|
||||
func (this *Context) GetInt(key string) (i int) { |
|
||||
if val, ok := this.Get(key); ok && val != nil { |
|
||||
i, _ = val.(int) |
|
||||
} |
|
||||
return |
|
||||
} |
|
||||
func (this *Context) GetInt64(key string) (i64 int64) { |
|
||||
if val, ok := this.Get(key); ok && val != nil { |
|
||||
i64, _ = val.(int64) |
|
||||
} |
|
||||
return |
|
||||
} |
|
||||
func (this *Context) GetUint(key string) (ui uint) { |
|
||||
if val, ok := this.Get(key); ok && val != nil { |
|
||||
ui, _ = val.(uint) |
|
||||
} |
|
||||
return |
|
||||
} |
|
||||
|
|
||||
func (c *Context) GetUInt32(key string) (i uint32) { |
|
||||
if val, ok := c.Get(key); ok && val != nil { |
|
||||
i, _ = val.(uint32) |
|
||||
} |
|
||||
return |
|
||||
} |
|
||||
|
|
||||
func (this *Context) GetUint64(key string) (ui64 uint64) { |
|
||||
if val, ok := this.Get(key); ok && val != nil { |
|
||||
ui64, _ = val.(uint64) |
|
||||
} |
|
||||
return |
|
||||
} |
|
||||
func (this *Context) GetFloat64(key string) (f64 float64) { |
|
||||
if val, ok := this.Get(key); ok && val != nil { |
|
||||
f64, _ = val.(float64) |
|
||||
} |
|
||||
return |
|
||||
} |
|
||||
func (this *Context) GetTime(key string) (t time.Time) { |
|
||||
if val, ok := this.Get(key); ok && val != nil { |
|
||||
t, _ = val.(time.Time) |
|
||||
} |
|
||||
return |
|
||||
} |
|
||||
func (this *Context) GetDuration(key string) (d time.Duration) { |
|
||||
if val, ok := this.Get(key); ok && val != nil { |
|
||||
d, _ = val.(time.Duration) |
|
||||
} |
|
||||
return |
|
||||
} |
|
||||
func (this *Context) GetStringSlice(key string) (ss []string) { |
|
||||
if val, ok := this.Get(key); ok && val != nil { |
|
||||
ss, _ = val.([]string) |
|
||||
} |
|
||||
return |
|
||||
} |
|
||||
func (this *Context) GetStringMap(key string) (sm map[string]interface{}) { |
|
||||
if val, ok := this.Get(key); ok && val != nil { |
|
||||
sm, _ = val.(map[string]interface{}) |
|
||||
} |
|
||||
return |
|
||||
} |
|
||||
func (this *Context) GetStringMapString(key string) (sms map[string]string) { |
|
||||
if val, ok := this.Get(key); ok && val != nil { |
|
||||
sms, _ = val.(map[string]string) |
|
||||
} |
|
||||
return |
|
||||
} |
|
||||
|
|
||||
func (this *Context) GetStringMapStringSlice(key string) (smss map[string][]string) { |
|
||||
if val, ok := this.Get(key); ok && val != nil { |
|
||||
smss, _ = val.(map[string][]string) |
|
||||
} |
|
||||
return |
|
||||
} |
|
||||
|
|
||||
func (this *Context) GetUserId() string { |
|
||||
return this.GetString("UserId") |
|
||||
} |
|
||||
|
|
||||
func (this *Context) Header(key, value string) { |
|
||||
if value == "" { |
|
||||
this.Writer.Header().Del(key) |
|
||||
return |
|
||||
} |
|
||||
this.Writer.Header().Set(key, value) |
|
||||
} |
|
||||
|
|
||||
// Status sets the HTTP response code.
|
|
||||
func (this *Context) Status(code int) { |
|
||||
this.Writer.WriteHeader(code) |
|
||||
} |
|
||||
|
|
||||
func (this *Context) Param(key string) string { |
|
||||
return this.Params.ByName(key) |
|
||||
} |
|
||||
|
|
||||
func (this *Context) AddParam(key, value string) { |
|
||||
this.Params = append(this.Params, Param{Key: key, Value: value}) |
|
||||
} |
|
||||
|
|
||||
func (this *Context) Query(key string) (value string) { |
|
||||
value, _ = this.GetQuery(key) |
|
||||
return |
|
||||
} |
|
||||
|
|
||||
func (this *Context) DefaultQuery(key, defaultValue string) string { |
|
||||
if value, ok := this.GetQuery(key); ok { |
|
||||
return value |
|
||||
} |
|
||||
return defaultValue |
|
||||
} |
|
||||
|
|
||||
func (this *Context) GetQuery(key string) (string, bool) { |
|
||||
if values, ok := this.GetQueryArray(key); ok { |
|
||||
return values[0], ok |
|
||||
} |
|
||||
return "", false |
|
||||
} |
|
||||
|
|
||||
func (this *Context) initQueryCache() { |
|
||||
if this.queryCache == nil { |
|
||||
if this.Request != nil { |
|
||||
this.queryCache = this.Request.URL.Query() |
|
||||
} else { |
|
||||
this.queryCache = url.Values{} |
|
||||
} |
|
||||
} |
|
||||
} |
|
||||
|
|
||||
func (this *Context) GetQueryArray(key string) (values []string, ok bool) { |
|
||||
this.initQueryCache() |
|
||||
values, ok = this.queryCache[key] |
|
||||
return |
|
||||
} |
|
||||
|
|
||||
func (this *Context) QueryMap(key string) (dicts map[string]string) { |
|
||||
dicts, _ = this.GetQueryMap(key) |
|
||||
return |
|
||||
} |
|
||||
|
|
||||
func (this *Context) GetQueryMap(key string) (map[string]string, bool) { |
|
||||
this.initQueryCache() |
|
||||
return this.get(this.queryCache, key) |
|
||||
} |
|
||||
|
|
||||
func (this *Context) PostForm(key string) (value string) { |
|
||||
value, _ = this.GetPostForm(key) |
|
||||
return |
|
||||
} |
|
||||
|
|
||||
func (this *Context) GetPostForm(key string) (string, bool) { |
|
||||
if values, ok := this.GetPostFormArray(key); ok { |
|
||||
return values[0], ok |
|
||||
} |
|
||||
return "", false |
|
||||
} |
|
||||
|
|
||||
func (this *Context) initFormCache() { |
|
||||
if this.formCache == nil { |
|
||||
this.formCache = make(url.Values) |
|
||||
req := this.Request |
|
||||
if err := req.ParseMultipartForm(this.engine.MaxMultipartMemory); err != nil { |
|
||||
if !errors.Is(err, http.ErrNotMultipart) { |
|
||||
this.Log.Errorf("error on parse multipart form array: %v", err) |
|
||||
} |
|
||||
} |
|
||||
this.formCache = req.PostForm |
|
||||
} |
|
||||
} |
|
||||
|
|
||||
func (this *Context) GetPostFormArray(key string) (values []string, ok bool) { |
|
||||
this.initFormCache() |
|
||||
values, ok = this.formCache[key] |
|
||||
return |
|
||||
} |
|
||||
|
|
||||
func (this *Context) PostFormMap(key string) (dicts map[string]string) { |
|
||||
dicts, _ = this.GetPostFormMap(key) |
|
||||
return |
|
||||
} |
|
||||
|
|
||||
func (this *Context) GetPostFormMap(key string) (map[string]string, bool) { |
|
||||
this.initFormCache() |
|
||||
return this.get(this.formCache, key) |
|
||||
} |
|
||||
|
|
||||
func (this *Context) FormFile(name string) (multipart.File, *multipart.FileHeader, error) { |
|
||||
if this.Request.MultipartForm == nil { |
|
||||
if err := this.Request.ParseMultipartForm(this.engine.MaxMultipartMemory); err != nil { |
|
||||
return nil, nil, err |
|
||||
} |
|
||||
} |
|
||||
f, fh, err := this.Request.FormFile(name) |
|
||||
if err != nil { |
|
||||
return nil, nil, err |
|
||||
} |
|
||||
return f, fh, err |
|
||||
} |
|
||||
|
|
||||
/* |
|
||||
MultipartForm 是解析后的多部分表单,包括文件上传。 |
|
||||
*/ |
|
||||
func (this *Context) MultipartForm() (*multipart.Form, error) { |
|
||||
err := this.Request.ParseMultipartForm(this.engine.MaxMultipartMemory) |
|
||||
return this.Request.MultipartForm, err |
|
||||
} |
|
||||
|
|
||||
/* |
|
||||
保存上传文件 |
|
||||
*/ |
|
||||
func (this *Context) SaveUploadedFile(file *multipart.FileHeader, dst string) error { |
|
||||
src, err := file.Open() |
|
||||
if err != nil { |
|
||||
return err |
|
||||
} |
|
||||
defer src.Close() |
|
||||
|
|
||||
out, err := os.Create(dst) |
|
||||
if err != nil { |
|
||||
return err |
|
||||
} |
|
||||
defer out.Close() |
|
||||
|
|
||||
_, err = io.Copy(out, src) |
|
||||
return err |
|
||||
} |
|
||||
func (this *Context) GetHeader(key string) string { |
|
||||
return this.requestHeader(key) |
|
||||
} |
|
||||
|
|
||||
func (this *Context) GetRawData() ([]byte, error) { |
|
||||
return ioutil.ReadAll(this.Request.Body) |
|
||||
} |
|
||||
|
|
||||
// 序列化--------------------------------------------------------------------------------------------
|
|
||||
func (this *Context) Bind(obj interface{}) error { |
|
||||
b := binding.Default(this.Request.Method, this.ContentType()) |
|
||||
return this.MustBindWith(obj, b) |
|
||||
} |
|
||||
|
|
||||
func (this *Context) ShouldBindJSON(obj interface{}) error { |
|
||||
return this.ShouldBindWith(obj, binding.JSON) |
|
||||
} |
|
||||
|
|
||||
func (this *Context) MustBindWith(obj interface{}, b binding.Binding) error { |
|
||||
if err := this.ShouldBindWith(obj, b); err != nil { |
|
||||
this.AbortWithError(http.StatusBadRequest, err).SetType(ErrorTypeBind) // nolint: errcheck
|
|
||||
return err |
|
||||
} |
|
||||
return nil |
|
||||
} |
|
||||
|
|
||||
func (this *Context) ShouldBindWith(obj interface{}, b binding.Binding) error { |
|
||||
return b.Bind(this.Request, obj) |
|
||||
} |
|
||||
func (this *Context) ShouldBindUri(obj interface{}) error { |
|
||||
m := make(map[string][]string) |
|
||||
for _, v := range this.Params { |
|
||||
m[v.Key] = []string{v.Value} |
|
||||
} |
|
||||
return binding.Uri.BindUri(m, obj) |
|
||||
} |
|
||||
|
|
||||
func (this *Context) BindJSON(obj interface{}) error { |
|
||||
return this.MustBindWith(obj, binding.JSON) |
|
||||
} |
|
||||
func (this *Context) BindXML(obj interface{}) error { |
|
||||
return this.MustBindWith(obj, binding.XML) |
|
||||
} |
|
||||
func (this *Context) BindQuery(obj interface{}) error { |
|
||||
return this.MustBindWith(obj, binding.Query) |
|
||||
} |
|
||||
|
|
||||
func (this *Context) BindYAML(obj interface{}) error { |
|
||||
return this.MustBindWith(obj, binding.YAML) |
|
||||
} |
|
||||
|
|
||||
func (this *Context) BindHeader(obj interface{}) error { |
|
||||
return this.MustBindWith(obj, binding.Header) |
|
||||
} |
|
||||
|
|
||||
func (this *Context) BindUri(obj interface{}) error { |
|
||||
if err := this.ShouldBindUri(obj); err != nil { |
|
||||
this.AbortWithError(http.StatusBadRequest, err).SetType(ErrorTypeBind) // nolint: errcheck
|
|
||||
return err |
|
||||
} |
|
||||
return nil |
|
||||
} |
|
||||
|
|
||||
// 输出-----------------------------------------------------------------------------------------
|
|
||||
func (this *Context) HTML(code int, name string, obj interface{}) { |
|
||||
instance := this.engine.HTMLRender.Instance(name, obj) |
|
||||
this.Render(code, instance) |
|
||||
} |
|
||||
|
|
||||
func (this *Context) IndentedJSON(code int, obj interface{}) { |
|
||||
this.Render(code, render.IndentedJSON{Data: obj}) |
|
||||
} |
|
||||
|
|
||||
func (this *Context) SecureJSON(code int, obj interface{}) { |
|
||||
this.Render(code, render.SecureJSON{Prefix: this.engine.secureJSONPrefix, Data: obj}) |
|
||||
} |
|
||||
|
|
||||
func (this *Context) JSONP(code int, obj interface{}) { |
|
||||
callback := this.DefaultQuery("callback", "") |
|
||||
if callback == "" { |
|
||||
this.Render(code, render.JSON{Data: obj}) |
|
||||
return |
|
||||
} |
|
||||
this.Render(code, render.JsonpJSON{Callback: callback, Data: obj}) |
|
||||
} |
|
||||
|
|
||||
func (this *Context) JSON(code int, obj interface{}) { |
|
||||
this.Render(code, render.JSON{Data: obj}) |
|
||||
} |
|
||||
|
|
||||
func (this *Context) AsciiJSON(code int, obj interface{}) { |
|
||||
this.Render(code, render.AsciiJSON{Data: obj}) |
|
||||
} |
|
||||
|
|
||||
func (this *Context) PureJSON(code int, obj interface{}) { |
|
||||
this.Render(code, render.PureJSON{Data: obj}) |
|
||||
} |
|
||||
|
|
||||
func (this *Context) XML(code int, obj interface{}) { |
|
||||
this.Render(code, render.XML{Data: obj}) |
|
||||
} |
|
||||
|
|
||||
func (this *Context) YAML(code int, obj interface{}) { |
|
||||
this.Render(code, render.YAML{Data: obj}) |
|
||||
} |
|
||||
|
|
||||
func (this *Context) ProtoBuf(code int, obj interface{}) { |
|
||||
this.Render(code, render.ProtoBuf{Data: obj}) |
|
||||
} |
|
||||
|
|
||||
func (this *Context) String(code int, format string, values ...interface{}) { |
|
||||
this.Render(code, render.String{Format: format, Data: values}) |
|
||||
} |
|
||||
|
|
||||
func (this *Context) Redirect(code int, location string) { |
|
||||
this.Render(-1, render.Redirect{ |
|
||||
Code: code, |
|
||||
Location: location, |
|
||||
Request: this.Request, |
|
||||
}) |
|
||||
} |
|
||||
|
|
||||
func (this *Context) Data(code int, contentType string, data []byte) { |
|
||||
this.Render(code, render.Data{ |
|
||||
ContentType: contentType, |
|
||||
Data: data, |
|
||||
}) |
|
||||
} |
|
||||
|
|
||||
func (this *Context) DataFromReader(code int, contentLength int64, contentType string, reader io.Reader, extraHeaders map[string]string) { |
|
||||
this.Render(code, render.Reader{ |
|
||||
Headers: extraHeaders, |
|
||||
ContentType: contentType, |
|
||||
ContentLength: contentLength, |
|
||||
Reader: reader, |
|
||||
}) |
|
||||
} |
|
||||
|
|
||||
func (this *Context) File(filepath string) { |
|
||||
http.ServeFile(this.Writer, this.Request, filepath) |
|
||||
} |
|
||||
|
|
||||
/*渲染页面接口*/ |
|
||||
func (this *Context) Render(code int, r render.Render) { |
|
||||
this.Status(code) |
|
||||
|
|
||||
if !bodyAllowedForStatus(code) { |
|
||||
r.WriteContentType(this.Writer) |
|
||||
this.Writer.WriteHeaderNow() |
|
||||
return |
|
||||
} |
|
||||
if err := r.Render(this.Writer); err != nil { |
|
||||
panic(err) |
|
||||
} |
|
||||
} |
|
||||
|
|
||||
/*渲染页面接口*/ |
|
||||
func (this *Context) RenderForBytes(code int, contentType string, body []byte) { |
|
||||
this.Status(code) |
|
||||
this.Header("ContentType", contentType) |
|
||||
if _, err := this.writermem.Write(body); err != nil { |
|
||||
panic(err) |
|
||||
} |
|
||||
} |
|
||||
|
|
||||
func (this *Context) FileFromFS(filepath string, fs http.FileSystem) { |
|
||||
defer func(old string) { |
|
||||
this.Request.URL.Path = old |
|
||||
}(this.Request.URL.Path) |
|
||||
this.Request.URL.Path = filepath |
|
||||
http.FileServer(fs).ServeHTTP(this.Writer, this.Request) |
|
||||
} |
|
||||
|
|
||||
/* |
|
||||
以高效的方式将指定的文件写入正文流 |
|
||||
在客户端,通常会使用给定的文件名下载文件 |
|
||||
*/ |
|
||||
func (this *Context) FileAttachment(filepath, filename string) { |
|
||||
if isASCII(filename) { |
|
||||
this.Writer.Header().Set("Content-Disposition", `attachment; filename="`+filename+`"`) |
|
||||
} else { |
|
||||
this.Writer.Header().Set("Content-Disposition", `attachment; filename*=UTF-8''`+url.QueryEscape(filename)) |
|
||||
} |
|
||||
http.ServeFile(this.Writer, this.Request, filepath) |
|
||||
} |
|
||||
|
|
||||
func (this *Context) Stream(step func(w io.Writer) bool) bool { |
|
||||
w := this.Writer |
|
||||
clientGone := w.CloseNotify() |
|
||||
for { |
|
||||
select { |
|
||||
case <-clientGone: |
|
||||
return true |
|
||||
default: |
|
||||
keepOpen := step(w) |
|
||||
w.Flush() |
|
||||
if !keepOpen { |
|
||||
return false |
|
||||
} |
|
||||
} |
|
||||
} |
|
||||
} |
|
||||
|
|
||||
type Negotiate struct { |
|
||||
Offered []string |
|
||||
HTMLName string |
|
||||
HTMLData interface{} |
|
||||
JSONData interface{} |
|
||||
XMLData interface{} |
|
||||
YAMLData interface{} |
|
||||
Data interface{} |
|
||||
} |
|
||||
|
|
||||
func (this *Context) Negotiate(code int, config Negotiate) { |
|
||||
switch this.NegotiateFormat(config.Offered...) { |
|
||||
case binding.MIMEJSON: |
|
||||
data := chooseData(config.JSONData, config.Data) |
|
||||
this.JSON(code, data) |
|
||||
|
|
||||
case binding.MIMEHTML: |
|
||||
data := chooseData(config.HTMLData, config.Data) |
|
||||
this.HTML(code, config.HTMLName, data) |
|
||||
|
|
||||
case binding.MIMEXML: |
|
||||
data := chooseData(config.XMLData, config.Data) |
|
||||
this.XML(code, data) |
|
||||
|
|
||||
case binding.MIMEYAML: |
|
||||
data := chooseData(config.YAMLData, config.Data) |
|
||||
this.YAML(code, data) |
|
||||
|
|
||||
default: |
|
||||
this.AbortWithError(http.StatusNotAcceptable, errors.New("the accepted formats are not offered by the server")) // nolint: errcheck
|
|
||||
} |
|
||||
} |
|
||||
|
|
||||
func (this *Context) NegotiateFormat(offered ...string) string { |
|
||||
assert1(len(offered) > 0, "you must provide at least one offer") |
|
||||
|
|
||||
if this.Accepted == nil { |
|
||||
this.Accepted = parseAccept(this.requestHeader("Accept")) |
|
||||
} |
|
||||
if len(this.Accepted) == 0 { |
|
||||
return offered[0] |
|
||||
} |
|
||||
for _, accepted := range this.Accepted { |
|
||||
for _, offer := range offered { |
|
||||
// According to RFC 2616 and RFC 2396, non-ASCII characters are not allowed in headers,
|
|
||||
// therefore we can just iterate over the string without casting it into []rune
|
|
||||
i := 0 |
|
||||
for ; i < len(accepted); i++ { |
|
||||
if accepted[i] == '*' || offer[i] == '*' { |
|
||||
return offer |
|
||||
} |
|
||||
if accepted[i] != offer[i] { |
|
||||
break |
|
||||
} |
|
||||
} |
|
||||
if i == len(accepted) { |
|
||||
return offer |
|
||||
} |
|
||||
} |
|
||||
} |
|
||||
return "" |
|
||||
} |
|
||||
|
|
||||
func (this *Context) SetAccepted(formats ...string) { |
|
||||
this.Accepted = formats |
|
||||
} |
|
||||
|
|
||||
func (this *Context) Deadline() (deadline time.Time, ok bool) { |
|
||||
if this.Request == nil || this.Request.Context() == nil { |
|
||||
return |
|
||||
} |
|
||||
return this.Request.Context().Deadline() |
|
||||
} |
|
||||
|
|
||||
func (this *Context) Done() <-chan struct{} { |
|
||||
if this.Request == nil || this.Request.Context() == nil { |
|
||||
return nil |
|
||||
} |
|
||||
return this.Request.Context().Done() |
|
||||
} |
|
||||
|
|
||||
func (this *Context) Err() error { |
|
||||
if this.Request == nil || this.Request.Context() == nil { |
|
||||
return nil |
|
||||
} |
|
||||
return this.Request.Context().Err() |
|
||||
} |
|
||||
|
|
||||
func (c *Context) Value(key interface{}) interface{} { |
|
||||
if key == 0 { |
|
||||
return c.Request |
|
||||
} |
|
||||
if keyAsString, ok := key.(string); ok { |
|
||||
if val, exists := c.Get(keyAsString); exists { |
|
||||
return val |
|
||||
} |
|
||||
} |
|
||||
if c.Request == nil || c.Request.Context() == nil { |
|
||||
return nil |
|
||||
} |
|
||||
return c.Request.Context().Value(key) |
|
||||
} |
|
||||
|
|
||||
func (this *Context) ContentType() string { |
|
||||
return filterFlags(this.requestHeader("Content-Type")) |
|
||||
} |
|
||||
|
|
||||
func (this *Context) RemoteIP() string { |
|
||||
ip, _, err := net.SplitHostPort(strings.TrimSpace(this.Request.RemoteAddr)) |
|
||||
if err != nil { |
|
||||
return "" |
|
||||
} |
|
||||
return ip |
|
||||
} |
|
||||
|
|
||||
func (this *Context) ClientIP() string { |
|
||||
// 检查我们是否在受信任的平台上运行,如果出错则继续向后运行
|
|
||||
if this.engine.TrustedPlatform != "" { |
|
||||
// Developers can define their own header of Trusted Platform or use predefined constants
|
|
||||
if addr := this.requestHeader(this.engine.TrustedPlatform); addr != "" { |
|
||||
return addr |
|
||||
} |
|
||||
} |
|
||||
|
|
||||
/* |
|
||||
// 它还检查 remoteIP 是否是受信任的代理。
|
|
||||
// 为了执行此验证,它将查看 IP 是否包含在至少一个 CIDR 块中
|
|
||||
// 由 Engine.SetTrustedProxies() 定义
|
|
||||
*/ |
|
||||
remoteIP := net.ParseIP(this.RemoteIP()) |
|
||||
if remoteIP == nil { |
|
||||
return "" |
|
||||
} |
|
||||
trusted := this.engine.isTrustedProxy(remoteIP) |
|
||||
|
|
||||
if trusted && this.engine.ForwardedByClientIP && this.engine.RemoteIPHeaders != nil { |
|
||||
for _, headerName := range this.engine.RemoteIPHeaders { |
|
||||
ip, valid := this.engine.validateHeader(this.requestHeader(headerName)) |
|
||||
if valid { |
|
||||
return ip |
|
||||
} |
|
||||
} |
|
||||
} |
|
||||
return remoteIP.String() |
|
||||
} |
|
||||
|
|
||||
func (this *Context) IsWebsocket() bool { |
|
||||
if strings.Contains(strings.ToLower(this.requestHeader("Connection")), "upgrade") && |
|
||||
strings.EqualFold(this.requestHeader("Upgrade"), "websocket") { |
|
||||
return true |
|
||||
} |
|
||||
return false |
|
||||
} |
|
||||
|
|
||||
func (this *Context) SetSameSite(samesite http.SameSite) { |
|
||||
this.sameSite = samesite |
|
||||
} |
|
||||
|
|
||||
func (this *Context) SetCookie(name, value string, maxAge int, path, domain string, secure, httpOnly bool) { |
|
||||
if path == "" { |
|
||||
path = "/" |
|
||||
} |
|
||||
http.SetCookie(this.Writer, &http.Cookie{ |
|
||||
Name: name, |
|
||||
Value: url.QueryEscape(value), |
|
||||
MaxAge: maxAge, |
|
||||
Path: path, |
|
||||
Domain: domain, |
|
||||
SameSite: this.sameSite, |
|
||||
Secure: secure, |
|
||||
HttpOnly: httpOnly, |
|
||||
}) |
|
||||
} |
|
||||
|
|
||||
func (this *Context) Cookie(name string) (string, error) { |
|
||||
cookie, err := this.Request.Cookie(name) |
|
||||
if err != nil { |
|
||||
return "", err |
|
||||
} |
|
||||
val, _ := url.QueryUnescape(cookie.Value) |
|
||||
return val, nil |
|
||||
} |
|
||||
|
|
||||
func (this *Context) Error(err error) *Error { |
|
||||
if err == nil { |
|
||||
panic("err is nil") |
|
||||
} |
|
||||
|
|
||||
var parsedError *Error |
|
||||
ok := errors.As(err, &parsedError) |
|
||||
if !ok { |
|
||||
parsedError = &Error{ |
|
||||
Err: err, |
|
||||
Type: ErrorTypePrivate, |
|
||||
} |
|
||||
} |
|
||||
|
|
||||
this.Errors = append(this.Errors, parsedError) |
|
||||
return parsedError |
|
||||
} |
|
||||
|
|
||||
func (this *Context) get(m map[string][]string, key string) (map[string]string, bool) { |
|
||||
dicts := make(map[string]string) |
|
||||
exist := false |
|
||||
for k, v := range m { |
|
||||
if i := strings.IndexByte(k, '['); i >= 1 && k[0:i] == key { |
|
||||
if j := strings.IndexByte(k[i+1:], ']'); j >= 1 { |
|
||||
exist = true |
|
||||
dicts[k[i+1:][:j]] = v[0] |
|
||||
} |
|
||||
} |
|
||||
} |
|
||||
return dicts, exist |
|
||||
} |
|
||||
|
|
||||
func (this *Context) requestHeader(key string) string { |
|
||||
return this.Request.Header.Get(key) |
|
||||
} |
|
||||
|
|
||||
func (this *Context) reset() { |
|
||||
this.Writer = &this.writermem |
|
||||
this.Params = this.Params[:0] |
|
||||
this.handlers = nil |
|
||||
this.index = -1 |
|
||||
|
|
||||
this.fullPath = "" |
|
||||
this.Keys = nil |
|
||||
this.Errors = this.Errors[:0] |
|
||||
this.Accepted = nil |
|
||||
this.queryCache = nil |
|
||||
this.formCache = nil |
|
||||
this.sameSite = 0 |
|
||||
*this.params = (*this.params)[:0] |
|
||||
*this.skippedNodes = (*this.skippedNodes)[:0] |
|
||||
} |
|
||||
|
|
||||
func (this *Context) ShouldBindQuery(obj any) error { |
|
||||
return this.ShouldBindWith(obj, binding.Query) |
|
||||
} |
|
||||
|
|
||||
/* |
|
||||
bodyAllowedForStatus 是 http.bodyAllowedForStatus 非导出函数的副本。 |
|
||||
*/ |
|
||||
func bodyAllowedForStatus(status int) bool { |
|
||||
switch { |
|
||||
case status >= 100 && status <= 199: |
|
||||
return false |
|
||||
case status == http.StatusNoContent: |
|
||||
return false |
|
||||
case status == http.StatusNotModified: |
|
||||
return false |
|
||||
} |
|
||||
return true |
|
||||
} |
|
||||
@ -1,44 +0,0 @@ |
|||||
package engine |
|
||||
|
|
||||
import ( |
|
||||
"net/http" |
|
||||
) |
|
||||
|
|
||||
type HandlerFunc func(*Context) |
|
||||
type HandlersChain []HandlerFunc |
|
||||
|
|
||||
func (c HandlersChain) Last() HandlerFunc { |
|
||||
if length := len(c); length > 0 { |
|
||||
return c[length-1] |
|
||||
} |
|
||||
return nil |
|
||||
} |
|
||||
|
|
||||
type RouteInfo struct { |
|
||||
Method string |
|
||||
Path string |
|
||||
Handler string |
|
||||
HandlerFunc HandlerFunc |
|
||||
} |
|
||||
|
|
||||
type RoutesInfo []RouteInfo |
|
||||
|
|
||||
type IRoutes interface { |
|
||||
Register(rcvr interface{}) |
|
||||
NoRoute(handlers ...HandlerFunc) |
|
||||
Group(relativePath string, handlers ...HandlerFunc) IRoutes |
|
||||
Use(...HandlerFunc) IRoutes |
|
||||
Handle(string, string, ...HandlerFunc) IRoutes |
|
||||
Any(string, ...HandlerFunc) IRoutes |
|
||||
GET(string, ...HandlerFunc) IRoutes |
|
||||
POST(string, ...HandlerFunc) IRoutes |
|
||||
DELETE(string, ...HandlerFunc) IRoutes |
|
||||
PATCH(string, ...HandlerFunc) IRoutes |
|
||||
PUT(string, ...HandlerFunc) IRoutes |
|
||||
OPTIONS(string, ...HandlerFunc) IRoutes |
|
||||
HEAD(string, ...HandlerFunc) IRoutes |
|
||||
StaticFile(string, string) IRoutes |
|
||||
StaticFileFS(string, string, http.FileSystem) IRoutes |
|
||||
Static(string, string) IRoutes |
|
||||
StaticFS(string, http.FileSystem) IRoutes |
|
||||
} |
|
||||
@ -1,522 +0,0 @@ |
|||||
package engine |
|
||||
|
|
||||
import ( |
|
||||
"html/template" |
|
||||
"net" |
|
||||
"net/http" |
|
||||
"path" |
|
||||
"strings" |
|
||||
"sync" |
|
||||
|
|
||||
"yunyan/lego/sys/gin/render" |
|
||||
"yunyan/lego/sys/log" |
|
||||
"yunyan/lego/utils" |
|
||||
|
|
||||
"golang.org/x/net/http2" |
|
||||
"golang.org/x/net/http2/h2c" |
|
||||
) |
|
||||
|
|
||||
/* |
|
||||
默认文件上传的最大尺寸 |
|
||||
*/ |
|
||||
const defaultMultipartMemory = 32 << 20 // 32 MB
|
|
||||
/* |
|
||||
默认可信代理 |
|
||||
*/ |
|
||||
var defaultTrustedCIDRs = []*net.IPNet{ |
|
||||
{ // 0.0.0.0/0 (IPv4)
|
|
||||
IP: net.IP{0x0, 0x0, 0x0, 0x0}, |
|
||||
Mask: net.IPMask{0x0, 0x0, 0x0, 0x0}, |
|
||||
}, |
|
||||
{ // ::/0 (IPv6)
|
|
||||
IP: net.IP{0x0, 0x0, 0x0, 0x0, 0x0, 0x0, 0x0, 0x0, 0x0, 0x0, 0x0, 0x0, 0x0, 0x0, 0x0, 0x0}, |
|
||||
Mask: net.IPMask{0x0, 0x0, 0x0, 0x0, 0x0, 0x0, 0x0, 0x0, 0x0, 0x0, 0x0, 0x0, 0x0, 0x0, 0x0, 0x0}, |
|
||||
}, |
|
||||
} |
|
||||
|
|
||||
func NewEngine(opts ...Option) (engine *Engine) { |
|
||||
option, _ := newOptions(opts...) |
|
||||
|
|
||||
engine = &Engine{ |
|
||||
RouterGroup: RouterGroup{ |
|
||||
Handlers: nil, |
|
||||
basePath: "/", |
|
||||
root: true, |
|
||||
}, |
|
||||
log: option.Log, |
|
||||
FuncMap: template.FuncMap{}, |
|
||||
RedirectTrailingSlash: option.RedirectTrailingSlash, |
|
||||
RedirectFixedPath: false, |
|
||||
HandleMethodNotAllowed: false, |
|
||||
ForwardedByClientIP: true, |
|
||||
RemoteIPHeaders: []string{"X-Forwarded-For", "X-Real-IP"}, |
|
||||
UseRawPath: false, |
|
||||
RemoveExtraSlash: false, |
|
||||
UnescapePathValues: true, |
|
||||
MaxMultipartMemory: option.MultipartMemory, |
|
||||
trees: make(methodTrees, 0, 9), |
|
||||
delims: render.Delims{Left: "{{", Right: "}}"}, |
|
||||
secureJSONPrefix: "while(1);", |
|
||||
trustedProxies: []string{"0.0.0.0/0"}, |
|
||||
trustedCIDRs: defaultTrustedCIDRs, |
|
||||
} |
|
||||
engine.RouterGroup.engine = engine |
|
||||
engine.pool.New = func() interface{} { |
|
||||
return engine.allocateContext() |
|
||||
} |
|
||||
return |
|
||||
} |
|
||||
|
|
||||
var ( |
|
||||
default404Body = []byte("404 page not found") |
|
||||
default405Body = []byte("405 method not allowed") |
|
||||
) |
|
||||
var mimePlain = []string{MIMEPlain} |
|
||||
|
|
||||
type Engine struct { |
|
||||
RouterGroup |
|
||||
log log.ILogger |
|
||||
UseRawPath bool |
|
||||
/* |
|
||||
如果启用,路由器尝试修复当前请求路径,如果没有 |
|
||||
如果没有 |
|
||||
已为其注册句柄。 |
|
||||
第一个多余的路径元素,如 ../ 或 // 被删除。
|
|
||||
之后路由器对清理后的路径进行不区分大小写的查找。 |
|
||||
如果可以找到该路由的句柄,则路由器进行重定向 |
|
||||
到正确的路径,GET 请求的状态码为 301,而 GET 请求的状态码为 307 |
|
||||
所有其他请求方法。 |
|
||||
例如 /FOO 和 /..//Foo 可以重定向到 /foo。
|
|
||||
RedirectTrailingSlash 与此选项无关。 |
|
||||
*/ |
|
||||
RedirectFixedPath bool |
|
||||
/* |
|
||||
如果为真,路径值将不转义 |
|
||||
*/ |
|
||||
UnescapePathValues bool |
|
||||
/* |
|
||||
如果当前路由无法匹配,但启用自动重定向 |
|
||||
带有(不带)尾部斜杠的路径的处理程序存在。 |
|
||||
例如,如果 /foo/ 被请求,但路由只存在于 /foo,则 |
|
||||
对于 GET 请求,客户端被重定向到 /foo,http 状态码为 301 |
|
||||
对于所有其他请求方法,则为 307。 |
|
||||
*/ |
|
||||
RedirectTrailingSlash bool //
|
|
||||
|
|
||||
/* |
|
||||
如果启用,路由器检查是否允许其他方法 |
|
||||
当前路由,如果当前请求无法路由。 |
|
||||
如果是这种情况,则使用“不允许的方法”回答请求 |
|
||||
和 HTTP 状态码 405。 |
|
||||
如果不允许其他方法,则将请求委托给 NotFound |
|
||||
处理程序。 |
|
||||
*/ |
|
||||
HandleMethodNotAllowed bool |
|
||||
/* |
|
||||
可以从 URL 中解析出一个参数,即使带有额外的斜杠。 |
|
||||
*/ |
|
||||
RemoveExtraSlash bool |
|
||||
/* |
|
||||
TrustedPlatform 如果设置为值 gin.Platform* 的常量,则信任由设置的标头 |
|
||||
那个平台,比如判断客户端IP |
|
||||
*/ |
|
||||
TrustedPlatform string |
|
||||
/* |
|
||||
ForwardedByClientIP 如果启用,客户端 IP 将从请求的标头中解析 |
|
||||
匹配存储在 `(*gin.Engine).RemoteIPHeaders` 中的那些。 如果没有 IP |
|
||||
fetched, 它回退到从获取的 IP |
|
||||
`(*gin.Context).Request.RemoteAddr`。 |
|
||||
ForwardedByClientIP 布尔值 |
|
||||
*/ |
|
||||
ForwardedByClientIP bool |
|
||||
/* |
|
||||
RemoteIPHeaders 用于获取客户端 IP 时的 headers 列表 |
|
||||
`(*gin.Engine).ForwardedByClientIP` 为 `true` 并且 |
|
||||
`(*gin.Context).Request.RemoteAddr` 被至少一个匹配 |
|
||||
由 `(*gin.Engine).SetTrustedProxies()` 定义的列表的网络来源。 |
|
||||
*/ |
|
||||
RemoteIPHeaders []string |
|
||||
/* |
|
||||
文件上传的最大尺寸 |
|
||||
*/ |
|
||||
MaxMultipartMemory int64 |
|
||||
/* |
|
||||
是否使用H2C |
|
||||
*/ |
|
||||
UseH2C bool |
|
||||
delims render.Delims |
|
||||
secureJSONPrefix string |
|
||||
HTMLRender render.HTMLRender |
|
||||
FuncMap template.FuncMap |
|
||||
noRoute HandlersChain |
|
||||
noMethod HandlersChain |
|
||||
allNoRoute HandlersChain |
|
||||
allNoMethod HandlersChain |
|
||||
pool sync.Pool |
|
||||
trees methodTrees |
|
||||
maxParams uint16 |
|
||||
maxSections uint16 |
|
||||
trustedProxies []string |
|
||||
trustedCIDRs []*net.IPNet |
|
||||
} |
|
||||
|
|
||||
func (this *Engine) ServeHTTP(w http.ResponseWriter, req *http.Request) { |
|
||||
c := this.pool.Get().(*Context) |
|
||||
c.writermem.reset(w) |
|
||||
c.Request = req |
|
||||
c.reset() |
|
||||
this.handleHTTPRequest(c) |
|
||||
this.pool.Put(c) |
|
||||
} |
|
||||
|
|
||||
func (this *Engine) Handler() http.Handler { |
|
||||
if !this.UseH2C { |
|
||||
return this |
|
||||
} |
|
||||
|
|
||||
h2s := &http2.Server{} |
|
||||
return h2c.NewHandler(this, h2s) |
|
||||
} |
|
||||
|
|
||||
/* |
|
||||
使用中间件 |
|
||||
*/ |
|
||||
func (this *Engine) Use(middleware ...HandlerFunc) IRoutes { |
|
||||
this.RouterGroup.Use(middleware...) |
|
||||
this.rebuild404Handlers() |
|
||||
this.rebuild405Handlers() |
|
||||
return this |
|
||||
} |
|
||||
|
|
||||
/* |
|
||||
LoadHTMLGlob 加载由 glob 模式标识的 HTML 文件 |
|
||||
并将结果与 HTML 渲染器相关联。 |
|
||||
*/ |
|
||||
func (this *Engine) LoadHTMLGlob(pattern string) { |
|
||||
left := this.delims.Left |
|
||||
right := this.delims.Right |
|
||||
templ := template.Must(template.New("").Delims(left, right).Funcs(this.FuncMap).ParseGlob(pattern)) |
|
||||
|
|
||||
if this.log.Enabled(log.DebugLevel) { |
|
||||
this.debugPrintLoadTemplate(templ) |
|
||||
this.HTMLRender = render.HTMLDebug{Glob: pattern, FuncMap: this.FuncMap, Delims: this.delims} |
|
||||
return |
|
||||
} |
|
||||
|
|
||||
this.SetHTMLTemplate(templ) |
|
||||
} |
|
||||
|
|
||||
/* |
|
||||
LoadHTMLFiles 加载一段 HTML 文件 |
|
||||
并将结果与 HTML 渲染器相关联。 |
|
||||
*/ |
|
||||
func (this *Engine) LoadHTMLFiles(files ...string) { |
|
||||
if this.log.Enabled(log.DebugLevel) { |
|
||||
this.HTMLRender = render.HTMLDebug{Files: files, FuncMap: this.FuncMap, Delims: this.delims} |
|
||||
return |
|
||||
} |
|
||||
templ := template.Must(template.New("").Delims(this.delims.Left, this.delims.Right).Funcs(this.FuncMap).ParseFiles(files...)) |
|
||||
this.SetHTMLTemplate(templ) |
|
||||
} |
|
||||
|
|
||||
func (this *Engine) SetHTMLTemplate(templ *template.Template) { |
|
||||
if len(this.trees) > 0 { |
|
||||
this.log.Warnf(`Since SetHTMLTemplate() is NOT thread-safe. It should only be called |
|
||||
at initialization. ie. before any route is registered or the router is listening in a socket: |
|
||||
router := gin.Default() |
|
||||
router.SetHTMLTemplate(template) // << good place
|
|
||||
`) |
|
||||
|
|
||||
} |
|
||||
this.HTMLRender = render.HTMLProduction{Template: templ.Funcs(this.FuncMap)} |
|
||||
} |
|
||||
|
|
||||
/* |
|
||||
设置template FuncMap |
|
||||
*/ |
|
||||
func (engine *Engine) SetFuncMap(funcMap template.FuncMap) { |
|
||||
engine.FuncMap = funcMap |
|
||||
} |
|
||||
|
|
||||
/* |
|
||||
404 处理路由 |
|
||||
*/ |
|
||||
func (this *Engine) NoRoute(handlers ...HandlerFunc) { |
|
||||
this.noRoute = handlers |
|
||||
this.rebuild404Handlers() |
|
||||
} |
|
||||
|
|
||||
/* |
|
||||
没有找到对应的方法 |
|
||||
*/ |
|
||||
func (this *Engine) NoMethod(handlers ...HandlerFunc) { |
|
||||
this.noMethod = handlers |
|
||||
this.rebuild405Handlers() |
|
||||
} |
|
||||
|
|
||||
func (engine *Engine) Routes() (routes RoutesInfo) { |
|
||||
for _, tree := range engine.trees { |
|
||||
routes = iterate("", tree.method, routes, tree.root) |
|
||||
} |
|
||||
return routes |
|
||||
} |
|
||||
|
|
||||
/* |
|
||||
设置信任代理 |
|
||||
*/ |
|
||||
func (this *Engine) SetTrustedProxies(trustedProxies []string) error { |
|
||||
this.trustedProxies = trustedProxies |
|
||||
return this.parseTrustedProxies() |
|
||||
} |
|
||||
|
|
||||
func (this *Engine) addRoute(method, path string, handlers HandlersChain) { |
|
||||
assert1(path[0] == '/', "path must begin with '/'") |
|
||||
assert1(method != "", "HTTP method can not be empty") |
|
||||
assert1(len(handlers) > 0, "there must be at least one handler") |
|
||||
if this.log.Enabled(log.DebugLevel) { |
|
||||
nuHandlers := len(handlers) |
|
||||
handlerName := nameOfFunction(handlers.Last()) |
|
||||
this.log.Debugf("%s:%s --> %s handlers:%d", method, path, handlerName, nuHandlers) |
|
||||
} |
|
||||
root := this.trees.get(method) |
|
||||
if root == nil { |
|
||||
root = new(node) |
|
||||
root.fullPath = "/" |
|
||||
this.trees = append(this.trees, methodTree{method: method, root: root}) |
|
||||
} |
|
||||
root.addRoute(path, handlers) |
|
||||
// Update maxParams
|
|
||||
if paramsCount := countParams(path); paramsCount > this.maxParams { |
|
||||
this.maxParams = paramsCount |
|
||||
} |
|
||||
|
|
||||
if sectionsCount := countSections(path); sectionsCount > this.maxSections { |
|
||||
this.maxSections = sectionsCount |
|
||||
} |
|
||||
} |
|
||||
|
|
||||
// 重定向
|
|
||||
func (this *Engine) HandleContext(c *Context) { |
|
||||
oldIndexValue := c.index |
|
||||
c.reset() |
|
||||
this.handleHTTPRequest(c) |
|
||||
c.index = oldIndexValue |
|
||||
} |
|
||||
|
|
||||
func (this *Engine) handleHTTPRequest(c *Context) { |
|
||||
httpMethod := c.Request.Method |
|
||||
rPath := c.Request.URL.Path |
|
||||
unescape := false |
|
||||
if this.UseRawPath && len(c.Request.URL.RawPath) > 0 { |
|
||||
rPath = c.Request.URL.RawPath |
|
||||
unescape = this.UnescapePathValues |
|
||||
} |
|
||||
if this.RemoveExtraSlash { |
|
||||
rPath = cleanPath(rPath) |
|
||||
} |
|
||||
t := this.trees |
|
||||
for i, tl := 0, len(t); i < tl; i++ { |
|
||||
if t[i].method != httpMethod { |
|
||||
continue |
|
||||
} |
|
||||
root := t[i].root |
|
||||
// Find route in tree
|
|
||||
value := root.getValue(rPath, c.params, c.skippedNodes, unescape) |
|
||||
if value.params != nil { |
|
||||
c.Params = *value.params |
|
||||
} |
|
||||
if value.handlers != nil { |
|
||||
c.handlers = value.handlers |
|
||||
c.fullPath = value.fullPath |
|
||||
c.Next() |
|
||||
c.writermem.WriteHeaderNow() |
|
||||
return |
|
||||
} |
|
||||
if httpMethod != http.MethodConnect && rPath != "/" { |
|
||||
if value.tsr && this.RedirectTrailingSlash { |
|
||||
this.redirectTrailingSlash(c) |
|
||||
return |
|
||||
} |
|
||||
if this.RedirectFixedPath && this.redirectFixedPath(c, root, this.RedirectFixedPath) { |
|
||||
return |
|
||||
} |
|
||||
} |
|
||||
break |
|
||||
} |
|
||||
|
|
||||
if this.HandleMethodNotAllowed { |
|
||||
for _, tree := range this.trees { |
|
||||
if tree.method == httpMethod { |
|
||||
continue |
|
||||
} |
|
||||
if value := tree.root.getValue(rPath, nil, c.skippedNodes, unescape); value.handlers != nil { |
|
||||
c.handlers = this.allNoMethod |
|
||||
this.serveError(c, http.StatusMethodNotAllowed, default405Body) |
|
||||
return |
|
||||
} |
|
||||
} |
|
||||
} |
|
||||
c.handlers = this.allNoRoute |
|
||||
this.serveError(c, http.StatusNotFound, default404Body) |
|
||||
} |
|
||||
|
|
||||
func (this *Engine) rebuild404Handlers() { |
|
||||
this.allNoRoute = this.combineHandlers(this.noRoute) |
|
||||
} |
|
||||
|
|
||||
func (this *Engine) rebuild405Handlers() { |
|
||||
this.allNoMethod = this.combineHandlers(this.noMethod) |
|
||||
} |
|
||||
|
|
||||
func (this *Engine) IsUnsafeTrustedProxies() bool { |
|
||||
return this.isTrustedProxy(net.ParseIP("0.0.0.0")) || this.isTrustedProxy(net.ParseIP("::")) |
|
||||
} |
|
||||
|
|
||||
// validateHeader 将解析 X-Forwarded-For 标头并返回受信任的客户端 IP 地址
|
|
||||
func (this *Engine) validateHeader(header string) (clientIP string, valid bool) { |
|
||||
if header == "" { |
|
||||
return "", false |
|
||||
} |
|
||||
items := strings.Split(header, ",") |
|
||||
for i := len(items) - 1; i >= 0; i-- { |
|
||||
ipStr := strings.TrimSpace(items[i]) |
|
||||
ip := net.ParseIP(ipStr) |
|
||||
if ip == nil { |
|
||||
break |
|
||||
} |
|
||||
|
|
||||
// X-Forwarded-For is appended by proxy
|
|
||||
// Check IPs in reverse order and stop when find untrusted proxy
|
|
||||
if (i == 0) || (!this.isTrustedProxy(ip)) { |
|
||||
return ipStr, true |
|
||||
} |
|
||||
} |
|
||||
return "", false |
|
||||
} |
|
||||
|
|
||||
// /目标Ip是否是可信
|
|
||||
func (this *Engine) isTrustedProxy(ip net.IP) bool { |
|
||||
if this.trustedCIDRs == nil { |
|
||||
return false |
|
||||
} |
|
||||
for _, cidr := range this.trustedCIDRs { |
|
||||
if cidr.Contains(ip) { |
|
||||
return true |
|
||||
} |
|
||||
} |
|
||||
return false |
|
||||
} |
|
||||
|
|
||||
func (this *Engine) serveError(c *Context, code int, defaultMessage []byte) { |
|
||||
c.writermem.status = code |
|
||||
c.Next() |
|
||||
if c.writermem.Written() { |
|
||||
return |
|
||||
} |
|
||||
if c.writermem.Status() == code { |
|
||||
c.writermem.Header()["Content-Type"] = mimePlain |
|
||||
_, err := c.Writer.Write(defaultMessage) |
|
||||
if err != nil { |
|
||||
this.log.Errorf("[SYS-Gin] cannot write message to writer during serve error: %v", err) |
|
||||
} |
|
||||
return |
|
||||
} |
|
||||
c.writermem.WriteHeaderNow() |
|
||||
} |
|
||||
|
|
||||
func (this *Engine) redirectFixedPath(c *Context, root *node, trailingSlash bool) bool { |
|
||||
req := c.Request |
|
||||
rPath := req.URL.Path |
|
||||
|
|
||||
if fixedPath, ok := root.findCaseInsensitivePath(cleanPath(rPath), trailingSlash); ok { |
|
||||
req.URL.Path = utils.BytesToString(fixedPath) |
|
||||
this.redirectRequest(c) |
|
||||
return true |
|
||||
} |
|
||||
return false |
|
||||
} |
|
||||
|
|
||||
func (this *Engine) redirectTrailingSlash(c *Context) { |
|
||||
req := c.Request |
|
||||
p := req.URL.Path |
|
||||
if prefix := path.Clean(c.Request.Header.Get("X-Forwarded-Prefix")); prefix != "." { |
|
||||
p = prefix + "/" + req.URL.Path |
|
||||
} |
|
||||
req.URL.Path = p + "/" |
|
||||
if length := len(p); length > 1 && p[length-1] == '/' { |
|
||||
req.URL.Path = p[:length-1] |
|
||||
} |
|
||||
this.redirectRequest(c) |
|
||||
} |
|
||||
|
|
||||
func (this *Engine) redirectRequest(c *Context) { |
|
||||
req := c.Request |
|
||||
rPath := req.URL.Path |
|
||||
rURL := req.URL.String() |
|
||||
|
|
||||
code := http.StatusMovedPermanently // Permanent redirect, request with GET method
|
|
||||
if req.Method != http.MethodGet { |
|
||||
code = http.StatusTemporaryRedirect |
|
||||
} |
|
||||
this.log.Debugf("redirecting request %d: %s --> %s", code, rPath, rURL) |
|
||||
http.Redirect(c.Writer, req, rURL, code) |
|
||||
c.writermem.WriteHeaderNow() |
|
||||
} |
|
||||
|
|
||||
func (this *Engine) parseTrustedProxies() error { |
|
||||
trustedCIDRs, err := this.prepareTrustedCIDRs() |
|
||||
this.trustedCIDRs = trustedCIDRs |
|
||||
return err |
|
||||
} |
|
||||
|
|
||||
func (this *Engine) prepareTrustedCIDRs() ([]*net.IPNet, error) { |
|
||||
if this.trustedProxies == nil { |
|
||||
return nil, nil |
|
||||
} |
|
||||
|
|
||||
cidr := make([]*net.IPNet, 0, len(this.trustedProxies)) |
|
||||
for _, trustedProxy := range this.trustedProxies { |
|
||||
if !strings.Contains(trustedProxy, "/") { |
|
||||
ip := parseIP(trustedProxy) |
|
||||
if ip == nil { |
|
||||
return cidr, &net.ParseError{Type: "IP address", Text: trustedProxy} |
|
||||
} |
|
||||
|
|
||||
switch len(ip) { |
|
||||
case net.IPv4len: |
|
||||
trustedProxy += "/32" |
|
||||
case net.IPv6len: |
|
||||
trustedProxy += "/128" |
|
||||
} |
|
||||
} |
|
||||
_, cidrNet, err := net.ParseCIDR(trustedProxy) |
|
||||
if err != nil { |
|
||||
return cidr, err |
|
||||
} |
|
||||
cidr = append(cidr, cidrNet) |
|
||||
} |
|
||||
return cidr, nil |
|
||||
} |
|
||||
|
|
||||
func (this *Engine) allocateContext() *Context { |
|
||||
v := make(Params, 0, this.maxParams) |
|
||||
skippedNodes := make([]skippedNode, 0, this.maxSections) |
|
||||
// return &Context{Log: this.log, engine: this, params: &v, skippedNodes: &skippedNodes, writermem: ResponseWriter{log: this.log}}
|
|
||||
return newContext(this.log, this, &v, &skippedNodes) |
|
||||
} |
|
||||
|
|
||||
// 日志接口-------------------------------------------------------------
|
|
||||
func (this *Engine) debugPrintLoadTemplate(tmpl *template.Template) { |
|
||||
var buf strings.Builder |
|
||||
for _, tmpl := range tmpl.Templates() { |
|
||||
buf.WriteString("\t- ") |
|
||||
buf.WriteString(tmpl.Name()) |
|
||||
buf.WriteString("\n") |
|
||||
} |
|
||||
format := "Loaded HTML Templates (%d): \n%s\n" |
|
||||
if !strings.HasSuffix(format, "\n") { |
|
||||
format += "\n" |
|
||||
} |
|
||||
this.log.Debugf(format, len(tmpl.Templates()), buf.String()) |
|
||||
|
|
||||
} |
|
||||
@ -1,201 +0,0 @@ |
|||||
// Copyright 2014 Manu Martinez-Almeida. All rights reserved.
|
|
||||
// Use of this source code is governed by a MIT style
|
|
||||
// license that can be found in the LICENSE file.
|
|
||||
|
|
||||
package engine |
|
||||
|
|
||||
import ( |
|
||||
"encoding/xml" |
|
||||
"fmt" |
|
||||
"reflect" |
|
||||
"strings" |
|
||||
|
|
||||
json "github.com/json-iterator/go" |
|
||||
) |
|
||||
|
|
||||
// H is a shortcut for map[string]interface{}
|
|
||||
type H map[string]interface{} |
|
||||
|
|
||||
// MarshalXML allows type H to be used with xml.Marshal.
|
|
||||
func (h H) MarshalXML(e *xml.Encoder, start xml.StartElement) error { |
|
||||
start.Name = xml.Name{ |
|
||||
Space: "", |
|
||||
Local: "map", |
|
||||
} |
|
||||
if err := e.EncodeToken(start); err != nil { |
|
||||
return err |
|
||||
} |
|
||||
for key, value := range h { |
|
||||
elem := xml.StartElement{ |
|
||||
Name: xml.Name{Space: "", Local: key}, |
|
||||
Attr: []xml.Attr{}, |
|
||||
} |
|
||||
if err := e.EncodeElement(value, elem); err != nil { |
|
||||
return err |
|
||||
} |
|
||||
} |
|
||||
|
|
||||
return e.EncodeToken(xml.EndElement{Name: start.Name}) |
|
||||
} |
|
||||
|
|
||||
// ErrorType is an unsigned 64-bit error code as defined in the gin spec.
|
|
||||
type ErrorType uint64 |
|
||||
|
|
||||
const ( |
|
||||
// ErrorTypeBind is used when Context.Bind() fails.
|
|
||||
ErrorTypeBind ErrorType = 1 << 63 |
|
||||
// ErrorTypeRender is used when Context.Render() fails.
|
|
||||
ErrorTypeRender ErrorType = 1 << 62 |
|
||||
// ErrorTypePrivate indicates a private error.
|
|
||||
ErrorTypePrivate ErrorType = 1 << 0 |
|
||||
// ErrorTypePublic indicates a public error.
|
|
||||
ErrorTypePublic ErrorType = 1 << 1 |
|
||||
// ErrorTypeAny indicates any other error.
|
|
||||
ErrorTypeAny ErrorType = 1<<64 - 1 |
|
||||
// ErrorTypeNu indicates any other error.
|
|
||||
ErrorTypeNu = 2 |
|
||||
) |
|
||||
|
|
||||
// Error represents a error's specification.
|
|
||||
type Error struct { |
|
||||
Err error |
|
||||
Type ErrorType |
|
||||
Meta interface{} |
|
||||
} |
|
||||
|
|
||||
type errorMsgs []*Error |
|
||||
|
|
||||
var _ error = &Error{} |
|
||||
|
|
||||
// SetType sets the error's type.
|
|
||||
func (msg *Error) SetType(flags ErrorType) *Error { |
|
||||
msg.Type = flags |
|
||||
return msg |
|
||||
} |
|
||||
|
|
||||
// SetMeta sets the error's meta data.
|
|
||||
func (msg *Error) SetMeta(data interface{}) *Error { |
|
||||
msg.Meta = data |
|
||||
return msg |
|
||||
} |
|
||||
|
|
||||
// JSON creates a properly formatted JSON
|
|
||||
func (msg *Error) JSON() interface{} { |
|
||||
jsonData := H{} |
|
||||
if msg.Meta != nil { |
|
||||
value := reflect.ValueOf(msg.Meta) |
|
||||
switch value.Kind() { |
|
||||
case reflect.Struct: |
|
||||
return msg.Meta |
|
||||
case reflect.Map: |
|
||||
for _, key := range value.MapKeys() { |
|
||||
jsonData[key.String()] = value.MapIndex(key).Interface() |
|
||||
} |
|
||||
default: |
|
||||
jsonData["meta"] = msg.Meta |
|
||||
} |
|
||||
} |
|
||||
if _, ok := jsonData["error"]; !ok { |
|
||||
jsonData["error"] = msg.Error() |
|
||||
} |
|
||||
return jsonData |
|
||||
} |
|
||||
|
|
||||
// MarshalJSON implements the json.Marshaller interface.
|
|
||||
func (msg *Error) MarshalJSON() ([]byte, error) { |
|
||||
return json.Marshal(msg.JSON()) |
|
||||
} |
|
||||
|
|
||||
// Error implements the error interface.
|
|
||||
func (msg Error) Error() string { |
|
||||
return msg.Err.Error() |
|
||||
} |
|
||||
|
|
||||
// IsType judges one error.
|
|
||||
func (msg *Error) IsType(flags ErrorType) bool { |
|
||||
return (msg.Type & flags) > 0 |
|
||||
} |
|
||||
|
|
||||
// Unwrap returns the wrapped error, to allow interoperability with errors.Is(), errors.As() and errors.Unwrap()
|
|
||||
func (msg *Error) Unwrap() error { |
|
||||
return msg.Err |
|
||||
} |
|
||||
|
|
||||
// ByType returns a readonly copy filtered the byte.
|
|
||||
// ie ByType(gin.ErrorTypePublic) returns a slice of errors with type=ErrorTypePublic.
|
|
||||
func (a errorMsgs) ByType(typ ErrorType) errorMsgs { |
|
||||
if len(a) == 0 { |
|
||||
return nil |
|
||||
} |
|
||||
if typ == ErrorTypeAny { |
|
||||
return a |
|
||||
} |
|
||||
var result errorMsgs |
|
||||
for _, msg := range a { |
|
||||
if msg.IsType(typ) { |
|
||||
result = append(result, msg) |
|
||||
} |
|
||||
} |
|
||||
return result |
|
||||
} |
|
||||
|
|
||||
// Last returns the last error in the slice. It returns nil if the array is empty.
|
|
||||
// Shortcut for errors[len(errors)-1].
|
|
||||
func (a errorMsgs) Last() *Error { |
|
||||
if length := len(a); length > 0 { |
|
||||
return a[length-1] |
|
||||
} |
|
||||
return nil |
|
||||
} |
|
||||
|
|
||||
// Errors returns an array with all the error messages.
|
|
||||
// Example:
|
|
||||
//
|
|
||||
// c.Error(errors.New("first"))
|
|
||||
// c.Error(errors.New("second"))
|
|
||||
// c.Error(errors.New("third"))
|
|
||||
// c.Errors.Errors() // == []string{"first", "second", "third"}
|
|
||||
func (a errorMsgs) Errors() []string { |
|
||||
if len(a) == 0 { |
|
||||
return nil |
|
||||
} |
|
||||
errorStrings := make([]string, len(a)) |
|
||||
for i, err := range a { |
|
||||
errorStrings[i] = err.Error() |
|
||||
} |
|
||||
return errorStrings |
|
||||
} |
|
||||
|
|
||||
func (a errorMsgs) JSON() interface{} { |
|
||||
switch length := len(a); length { |
|
||||
case 0: |
|
||||
return nil |
|
||||
case 1: |
|
||||
return a.Last().JSON() |
|
||||
default: |
|
||||
jsonData := make([]interface{}, length) |
|
||||
for i, err := range a { |
|
||||
jsonData[i] = err.JSON() |
|
||||
} |
|
||||
return jsonData |
|
||||
} |
|
||||
} |
|
||||
|
|
||||
// MarshalJSON implements the json.Marshaller interface.
|
|
||||
func (a errorMsgs) MarshalJSON() ([]byte, error) { |
|
||||
return json.Marshal(a.JSON()) |
|
||||
} |
|
||||
|
|
||||
func (a errorMsgs) String() string { |
|
||||
if len(a) == 0 { |
|
||||
return "" |
|
||||
} |
|
||||
var buffer strings.Builder |
|
||||
for i, msg := range a { |
|
||||
fmt.Fprintf(&buffer, "Error #%02d: %s\n", i+1, msg.Err) |
|
||||
if msg.Meta != nil { |
|
||||
fmt.Fprintf(&buffer, " Meta: %v\n", msg.Meta) |
|
||||
} |
|
||||
} |
|
||||
return buffer.String() |
|
||||
} |
|
||||
@ -1,26 +0,0 @@ |
|||||
package engine |
|
||||
|
|
||||
import "net/http" |
|
||||
|
|
||||
type onlyFilesFS struct { |
|
||||
fs http.FileSystem |
|
||||
} |
|
||||
type neuteredReaddirFile struct { |
|
||||
http.File |
|
||||
} |
|
||||
|
|
||||
func Dir(root string, listDirectory bool) http.FileSystem { |
|
||||
fs := http.Dir(root) |
|
||||
if listDirectory { |
|
||||
return fs |
|
||||
} |
|
||||
return &onlyFilesFS{fs} |
|
||||
} |
|
||||
|
|
||||
func (fs onlyFilesFS) Open(name string) (http.File, error) { |
|
||||
f, err := fs.fs.Open(name) |
|
||||
if err != nil { |
|
||||
return nil, err |
|
||||
} |
|
||||
return neuteredReaddirFile{f}, nil |
|
||||
} |
|
||||
@ -1,47 +0,0 @@ |
|||||
package engine |
|
||||
|
|
||||
import ( |
|
||||
"yunyan/lego/sys/log" |
|
||||
) |
|
||||
|
|
||||
type Option func(*Options) |
|
||||
type Options struct { |
|
||||
MultipartMemory int64 //文件上传最大尺寸
|
|
||||
RedirectTrailingSlash bool |
|
||||
Log log.ILogger |
|
||||
} |
|
||||
|
|
||||
func SetMultipartMemory(v int64) Option { |
|
||||
return func(o *Options) { |
|
||||
o.MultipartMemory = v |
|
||||
} |
|
||||
} |
|
||||
|
|
||||
/* |
|
||||
如果当前路由无法匹配,但启用自动重定向 |
|
||||
带有(不带)尾部斜杠的路径的处理程序存在。 |
|
||||
例如,如果 /foo/ 被请求,但路由只存在于 /foo,则 |
|
||||
对于 GET 请求,客户端被重定向到 /foo,http 状态码为 301 |
|
||||
对于所有其他请求方法,则为 307。 |
|
||||
*/ |
|
||||
func SetRedirectTrailingSlash(v bool) Option { |
|
||||
return func(o *Options) { |
|
||||
o.RedirectTrailingSlash = v |
|
||||
} |
|
||||
} |
|
||||
|
|
||||
func SetLog(log log.ILogger) Option { |
|
||||
return func(o *Options) { |
|
||||
o.Log = log |
|
||||
} |
|
||||
} |
|
||||
|
|
||||
func newOptions(opts ...Option) (options *Options, err error) { |
|
||||
options = &Options{ |
|
||||
MultipartMemory: defaultMultipartMemory, |
|
||||
} |
|
||||
for _, o := range opts { |
|
||||
o(options) |
|
||||
} |
|
||||
return |
|
||||
} |
|
||||
@ -1,116 +0,0 @@ |
|||||
package engine |
|
||||
|
|
||||
import ( |
|
||||
"bufio" |
|
||||
"io" |
|
||||
"net" |
|
||||
"net/http" |
|
||||
|
|
||||
"yunyan/lego/sys/log" |
|
||||
) |
|
||||
|
|
||||
const ( |
|
||||
noWritten = -1 |
|
||||
defaultStatus = http.StatusOK |
|
||||
) |
|
||||
|
|
||||
type IResponseWriter interface { |
|
||||
http.ResponseWriter |
|
||||
http.Hijacker |
|
||||
http.Flusher |
|
||||
http.CloseNotifier |
|
||||
/* |
|
||||
返回当前请求的 HTTP 响应状态码。 |
|
||||
*/ |
|
||||
Status() int |
|
||||
/* |
|
||||
大小返回已经写入响应 http 正文的字节数 |
|
||||
*/ |
|
||||
Size() int |
|
||||
// WriteString writes the string into the response body.
|
|
||||
WriteString(string) (int, error) |
|
||||
// Written returns true if the response body was already written.
|
|
||||
Written() bool |
|
||||
/* |
|
||||
WriteHeaderNow 强制写入 http 标头(状态码 + 标头)。 |
|
||||
*/ |
|
||||
WriteHeaderNow() |
|
||||
// Pusher get the http.Pusher for server push
|
|
||||
Pusher() http.Pusher |
|
||||
} |
|
||||
type ResponseWriter struct { |
|
||||
http.ResponseWriter |
|
||||
log log.ILogger |
|
||||
size int |
|
||||
status int |
|
||||
} |
|
||||
|
|
||||
func (this *ResponseWriter) reset(writer http.ResponseWriter) { |
|
||||
this.ResponseWriter = writer |
|
||||
this.size = noWritten |
|
||||
this.status = defaultStatus |
|
||||
} |
|
||||
|
|
||||
func (this *ResponseWriter) WriteHeader(code int) { |
|
||||
if code > 0 && this.status != code { |
|
||||
if this.Written() { |
|
||||
this.log.Warnf("Headers were already written. Wanted to override status code %d with %d", this.status, code) |
|
||||
} |
|
||||
this.status = code |
|
||||
} |
|
||||
} |
|
||||
func (this *ResponseWriter) WriteHeaderNow() { |
|
||||
if !this.Written() { |
|
||||
this.size = 0 |
|
||||
this.ResponseWriter.WriteHeader(this.status) |
|
||||
} |
|
||||
} |
|
||||
|
|
||||
func (this *ResponseWriter) Write(data []byte) (n int, err error) { |
|
||||
this.WriteHeaderNow() |
|
||||
n, err = this.ResponseWriter.Write(data) |
|
||||
this.size += n |
|
||||
return |
|
||||
} |
|
||||
|
|
||||
func (this *ResponseWriter) WriteString(s string) (n int, err error) { |
|
||||
this.WriteHeaderNow() |
|
||||
n, err = io.WriteString(this.ResponseWriter, s) |
|
||||
this.size += n |
|
||||
return |
|
||||
} |
|
||||
|
|
||||
func (this *ResponseWriter) Status() int { |
|
||||
return this.status |
|
||||
} |
|
||||
|
|
||||
func (this *ResponseWriter) Size() int { |
|
||||
return this.size |
|
||||
} |
|
||||
func (this *ResponseWriter) Written() bool { |
|
||||
return this.size != noWritten |
|
||||
} |
|
||||
|
|
||||
// Hijack implements the http.Hijacker interface.
|
|
||||
func (this *ResponseWriter) Hijack() (net.Conn, *bufio.ReadWriter, error) { |
|
||||
if this.size < 0 { |
|
||||
this.size = 0 |
|
||||
} |
|
||||
return this.ResponseWriter.(http.Hijacker).Hijack() |
|
||||
} |
|
||||
|
|
||||
func (this *ResponseWriter) CloseNotify() <-chan bool { |
|
||||
return this.ResponseWriter.(http.CloseNotifier).CloseNotify() |
|
||||
} |
|
||||
|
|
||||
func (this *ResponseWriter) Flush() { |
|
||||
this.WriteHeaderNow() |
|
||||
this.ResponseWriter.(http.Flusher).Flush() |
|
||||
} |
|
||||
|
|
||||
func (this *ResponseWriter) Pusher() (pusher http.Pusher) { |
|
||||
if pusher, ok := this.ResponseWriter.(http.Pusher); ok { |
|
||||
return pusher |
|
||||
} |
|
||||
return nil |
|
||||
} |
|
||||
@ -1,196 +0,0 @@ |
|||||
package engine |
|
||||
|
|
||||
import ( |
|
||||
"net/http" |
|
||||
"path" |
|
||||
"reflect" |
|
||||
"regexp" |
|
||||
"strings" |
|
||||
) |
|
||||
|
|
||||
var ( |
|
||||
regEnLetter = regexp.MustCompile("^[A-Z]+$") |
|
||||
anyMethods = []string{ |
|
||||
http.MethodGet, http.MethodPost, http.MethodPut, http.MethodPatch, |
|
||||
http.MethodHead, http.MethodOptions, http.MethodDelete, http.MethodConnect, |
|
||||
http.MethodTrace, |
|
||||
} |
|
||||
) |
|
||||
|
|
||||
type RouterGroup struct { |
|
||||
Handlers HandlersChain |
|
||||
basePath string |
|
||||
engine *Engine |
|
||||
root bool |
|
||||
} |
|
||||
|
|
||||
func (this *RouterGroup) BasePath() string { |
|
||||
return this.basePath |
|
||||
} |
|
||||
func (this *RouterGroup) NoRoute(handlers ...HandlerFunc) { |
|
||||
this.engine.NoRoute(handlers...) |
|
||||
} |
|
||||
func (this *RouterGroup) Use(middleware ...HandlerFunc) IRoutes { |
|
||||
this.Handlers = append(this.Handlers, middleware...) |
|
||||
return this.returnObj() |
|
||||
} |
|
||||
|
|
||||
func (this *RouterGroup) Group(relativePath string, handlers ...HandlerFunc) IRoutes { |
|
||||
return &RouterGroup{ |
|
||||
Handlers: this.combineHandlers(handlers), |
|
||||
basePath: this.calculateAbsolutePath(relativePath), |
|
||||
engine: this.engine, |
|
||||
} |
|
||||
} |
|
||||
|
|
||||
func (this *RouterGroup) Handle(httpMethod, relativePath string, handlers ...HandlerFunc) IRoutes { |
|
||||
if matched := regEnLetter.MatchString(httpMethod); !matched { |
|
||||
panic("http method " + httpMethod + " is not valid") |
|
||||
} |
|
||||
return this.handle(httpMethod, relativePath, handlers) |
|
||||
} |
|
||||
|
|
||||
func (this *RouterGroup) POST(relativePath string, handlers ...HandlerFunc) IRoutes { |
|
||||
return this.handle(http.MethodPost, relativePath, handlers) |
|
||||
} |
|
||||
|
|
||||
func (this *RouterGroup) GET(relativePath string, handlers ...HandlerFunc) IRoutes { |
|
||||
return this.handle(http.MethodGet, relativePath, handlers) |
|
||||
} |
|
||||
|
|
||||
func (this *RouterGroup) DELETE(relativePath string, handlers ...HandlerFunc) IRoutes { |
|
||||
return this.handle(http.MethodDelete, relativePath, handlers) |
|
||||
} |
|
||||
|
|
||||
func (this *RouterGroup) PATCH(relativePath string, handlers ...HandlerFunc) IRoutes { |
|
||||
return this.handle(http.MethodPatch, relativePath, handlers) |
|
||||
} |
|
||||
|
|
||||
func (this *RouterGroup) PUT(relativePath string, handlers ...HandlerFunc) IRoutes { |
|
||||
return this.handle(http.MethodPut, relativePath, handlers) |
|
||||
} |
|
||||
|
|
||||
func (this *RouterGroup) OPTIONS(relativePath string, handlers ...HandlerFunc) IRoutes { |
|
||||
return this.handle(http.MethodOptions, relativePath, handlers) |
|
||||
} |
|
||||
|
|
||||
func (this *RouterGroup) HEAD(relativePath string, handlers ...HandlerFunc) IRoutes { |
|
||||
return this.handle(http.MethodHead, relativePath, handlers) |
|
||||
} |
|
||||
|
|
||||
func (this *RouterGroup) Any(relativePath string, handlers ...HandlerFunc) IRoutes { |
|
||||
for _, method := range anyMethods { |
|
||||
this.handle(method, relativePath, handlers) |
|
||||
} |
|
||||
return this.returnObj() |
|
||||
} |
|
||||
|
|
||||
func (this *RouterGroup) StaticFile(relativePath, filepath string) IRoutes { |
|
||||
return this.staticFileHandler(relativePath, func(c *Context) { |
|
||||
c.File(filepath) |
|
||||
}) |
|
||||
} |
|
||||
|
|
||||
func (this *RouterGroup) StaticFileFS(relativePath, filepath string, fs http.FileSystem) IRoutes { |
|
||||
return this.staticFileHandler(relativePath, func(c *Context) { |
|
||||
c.FileFromFS(filepath, fs) |
|
||||
}) |
|
||||
} |
|
||||
|
|
||||
func (this *RouterGroup) Static(relativePath, root string) IRoutes { |
|
||||
return this.StaticFS(relativePath, Dir(root, false)) |
|
||||
} |
|
||||
|
|
||||
func (this *RouterGroup) StaticFS(relativePath string, fs http.FileSystem) IRoutes { |
|
||||
if strings.Contains(relativePath, ":") || strings.Contains(relativePath, "*") { |
|
||||
panic("URL parameters can not be used when serving a static folder") |
|
||||
} |
|
||||
handler := this.createStaticHandler(relativePath, fs) |
|
||||
urlPattern := path.Join(relativePath, "/*filepath") |
|
||||
|
|
||||
// Register GET and HEAD handlers
|
|
||||
this.GET(urlPattern, handler) |
|
||||
this.HEAD(urlPattern, handler) |
|
||||
return this.returnObj() |
|
||||
} |
|
||||
|
|
||||
func (this *RouterGroup) handle(httpMethod, relativePath string, handlers HandlersChain) IRoutes { |
|
||||
absolutePath := this.calculateAbsolutePath(relativePath) |
|
||||
handlers = this.combineHandlers(handlers) |
|
||||
this.engine.addRoute(httpMethod, absolutePath, handlers) |
|
||||
return this.returnObj() |
|
||||
} |
|
||||
|
|
||||
func (group *RouterGroup) createStaticHandler(relativePath string, fs http.FileSystem) HandlerFunc { |
|
||||
absolutePath := group.calculateAbsolutePath(relativePath) |
|
||||
fileServer := http.StripPrefix(absolutePath, http.FileServer(fs)) |
|
||||
return func(c *Context) { |
|
||||
if _, noListing := fs.(*onlyFilesFS); noListing { |
|
||||
c.Writer.WriteHeader(http.StatusNotFound) |
|
||||
} |
|
||||
file := c.Param("filepath") |
|
||||
f, err := fs.Open(file) |
|
||||
if err != nil { |
|
||||
c.Writer.WriteHeader(http.StatusNotFound) |
|
||||
c.handlers = group.engine.noRoute |
|
||||
c.index = -1 |
|
||||
return |
|
||||
} |
|
||||
f.Close() |
|
||||
|
|
||||
fileServer.ServeHTTP(c.Writer, c.Request) |
|
||||
} |
|
||||
} |
|
||||
|
|
||||
func (this *RouterGroup) combineHandlers(handlers HandlersChain) HandlersChain { |
|
||||
finalSize := len(this.Handlers) + len(handlers) |
|
||||
assert1(finalSize < int(abortIndex), "too many handlers") |
|
||||
mergedHandlers := make(HandlersChain, finalSize) |
|
||||
copy(mergedHandlers, this.Handlers) |
|
||||
copy(mergedHandlers[len(this.Handlers):], handlers) |
|
||||
return mergedHandlers |
|
||||
} |
|
||||
|
|
||||
func (this *RouterGroup) calculateAbsolutePath(relativePath string) string { |
|
||||
return joinPaths(this.basePath, relativePath) |
|
||||
} |
|
||||
|
|
||||
func (this *RouterGroup) staticFileHandler(relativePath string, handler HandlerFunc) IRoutes { |
|
||||
if strings.Contains(relativePath, ":") || strings.Contains(relativePath, "*") { |
|
||||
panic("URL parameters can not be used when serving a static file") |
|
||||
} |
|
||||
this.GET(relativePath, handler) |
|
||||
this.HEAD(relativePath, handler) |
|
||||
return this.returnObj() |
|
||||
} |
|
||||
|
|
||||
func (this *RouterGroup) returnObj() IRoutes { |
|
||||
if this.root { |
|
||||
return this.engine |
|
||||
} |
|
||||
return this |
|
||||
} |
|
||||
|
|
||||
func (this *RouterGroup) Register(rcvr interface{}) { |
|
||||
typ := reflect.TypeOf(rcvr) |
|
||||
vof := reflect.ValueOf(rcvr) |
|
||||
for m := 0; m < typ.NumMethod(); m++ { |
|
||||
method := typ.Method(m) |
|
||||
mname := method.Name |
|
||||
mtype := method.Type |
|
||||
if method.PkgPath != "" { |
|
||||
continue |
|
||||
} |
|
||||
if mtype.NumIn() != 2 { |
|
||||
continue |
|
||||
} |
|
||||
context := mtype.In(1) |
|
||||
if context.String() != "*engine.Context" { |
|
||||
continue |
|
||||
} |
|
||||
if mtype.NumOut() != 0 { |
|
||||
continue |
|
||||
} |
|
||||
this.POST(strings.ToLower(mname), vof.MethodByName(mname).Interface().(func(*Context))) |
|
||||
} |
|
||||
} |
|
||||
@ -1,838 +0,0 @@ |
|||||
package engine |
|
||||
|
|
||||
import ( |
|
||||
"bytes" |
|
||||
"net/url" |
|
||||
"strings" |
|
||||
"unicode" |
|
||||
"unicode/utf8" |
|
||||
|
|
||||
"yunyan/lego/utils" |
|
||||
) |
|
||||
|
|
||||
var ( |
|
||||
strColon = []byte(":") |
|
||||
strStar = []byte("*") |
|
||||
strSlash = []byte("/") |
|
||||
) |
|
||||
|
|
||||
type Param struct { |
|
||||
Key string |
|
||||
Value string |
|
||||
} |
|
||||
|
|
||||
type Params []Param |
|
||||
|
|
||||
func (ps Params) Get(name string) (string, bool) { |
|
||||
for _, entry := range ps { |
|
||||
if entry.Key == name { |
|
||||
return entry.Value, true |
|
||||
} |
|
||||
} |
|
||||
return "", false |
|
||||
} |
|
||||
|
|
||||
func (ps Params) ByName(name string) (va string) { |
|
||||
va, _ = ps.Get(name) |
|
||||
return |
|
||||
} |
|
||||
|
|
||||
type nodeType uint8 |
|
||||
type nodeValue struct { |
|
||||
handlers HandlersChain |
|
||||
params *Params |
|
||||
tsr bool |
|
||||
fullPath string |
|
||||
} |
|
||||
type skippedNode struct { |
|
||||
path string |
|
||||
node *node |
|
||||
paramsCount int16 |
|
||||
} |
|
||||
|
|
||||
const ( |
|
||||
root nodeType = iota + 1 |
|
||||
param |
|
||||
catchAll |
|
||||
) |
|
||||
|
|
||||
type node struct { |
|
||||
path string |
|
||||
indices string |
|
||||
wildChild bool |
|
||||
nType nodeType |
|
||||
priority uint32 |
|
||||
children []*node // child nodes, at most 1 :param style node at the end of the array
|
|
||||
handlers HandlersChain |
|
||||
fullPath string |
|
||||
} |
|
||||
|
|
||||
func (n *node) addChild(child *node) { |
|
||||
if n.wildChild && len(n.children) > 0 { |
|
||||
wildcardChild := n.children[len(n.children)-1] |
|
||||
n.children = append(n.children[:len(n.children)-1], child, wildcardChild) |
|
||||
} else { |
|
||||
n.children = append(n.children, child) |
|
||||
} |
|
||||
} |
|
||||
|
|
||||
func (n *node) addRoute(path string, handlers HandlersChain) { |
|
||||
fullPath := path |
|
||||
n.priority++ |
|
||||
// Empty tree
|
|
||||
if len(n.path) == 0 && len(n.children) == 0 { |
|
||||
n.insertChild(path, fullPath, handlers) |
|
||||
n.nType = root |
|
||||
return |
|
||||
} |
|
||||
parentFullPathIndex := 0 |
|
||||
walk: |
|
||||
for { |
|
||||
// Find the longest common prefix.
|
|
||||
// This also implies that the common prefix contains no ':' or '*'
|
|
||||
// since the existing key can't contain those chars.
|
|
||||
i := longestCommonPrefix(path, n.path) |
|
||||
|
|
||||
// Split edge
|
|
||||
if i < len(n.path) { |
|
||||
child := node{ |
|
||||
path: n.path[i:], |
|
||||
wildChild: n.wildChild, |
|
||||
indices: n.indices, |
|
||||
children: n.children, |
|
||||
handlers: n.handlers, |
|
||||
priority: n.priority - 1, |
|
||||
fullPath: n.fullPath, |
|
||||
} |
|
||||
|
|
||||
n.children = []*node{&child} |
|
||||
// []byte for proper unicode char conversion, see #65
|
|
||||
n.indices = utils.BytesToString([]byte{n.path[i]}) |
|
||||
n.path = path[:i] |
|
||||
n.handlers = nil |
|
||||
n.wildChild = false |
|
||||
n.fullPath = fullPath[:parentFullPathIndex+i] |
|
||||
} |
|
||||
|
|
||||
// Make new node a child of this node
|
|
||||
if i < len(path) { |
|
||||
path = path[i:] |
|
||||
c := path[0] |
|
||||
|
|
||||
// '/' after param
|
|
||||
if n.nType == param && c == '/' && len(n.children) == 1 { |
|
||||
parentFullPathIndex += len(n.path) |
|
||||
n = n.children[0] |
|
||||
n.priority++ |
|
||||
continue walk |
|
||||
} |
|
||||
|
|
||||
// Check if a child with the next path byte exists
|
|
||||
for i, max := 0, len(n.indices); i < max; i++ { |
|
||||
if c == n.indices[i] { |
|
||||
parentFullPathIndex += len(n.path) |
|
||||
i = n.incrementChildPrio(i) |
|
||||
n = n.children[i] |
|
||||
continue walk |
|
||||
} |
|
||||
} |
|
||||
|
|
||||
// Otherwise insert it
|
|
||||
if c != ':' && c != '*' && n.nType != catchAll { |
|
||||
// []byte for proper unicode char conversion, see #65
|
|
||||
n.indices += utils.BytesToString([]byte{c}) |
|
||||
child := &node{ |
|
||||
fullPath: fullPath, |
|
||||
} |
|
||||
n.addChild(child) |
|
||||
n.incrementChildPrio(len(n.indices) - 1) |
|
||||
n = child |
|
||||
} else if n.wildChild { |
|
||||
// inserting a wildcard node, need to check if it conflicts with the existing wildcard
|
|
||||
n = n.children[len(n.children)-1] |
|
||||
n.priority++ |
|
||||
|
|
||||
// Check if the wildcard matches
|
|
||||
if len(path) >= len(n.path) && n.path == path[:len(n.path)] && |
|
||||
// Adding a child to a catchAll is not possible
|
|
||||
n.nType != catchAll && |
|
||||
// Check for longer wildcard, e.g. :name and :names
|
|
||||
(len(n.path) >= len(path) || path[len(n.path)] == '/') { |
|
||||
continue walk |
|
||||
} |
|
||||
|
|
||||
// Wildcard conflict
|
|
||||
pathSeg := path |
|
||||
if n.nType != catchAll { |
|
||||
pathSeg = strings.SplitN(pathSeg, "/", 2)[0] |
|
||||
} |
|
||||
prefix := fullPath[:strings.Index(fullPath, pathSeg)] + n.path |
|
||||
panic("'" + pathSeg + |
|
||||
"' in new path '" + fullPath + |
|
||||
"' conflicts with existing wildcard '" + n.path + |
|
||||
"' in existing prefix '" + prefix + |
|
||||
"'") |
|
||||
} |
|
||||
|
|
||||
n.insertChild(path, fullPath, handlers) |
|
||||
return |
|
||||
} |
|
||||
|
|
||||
// Otherwise add handle to current node
|
|
||||
if n.handlers != nil { |
|
||||
panic("handlers are already registered for path '" + fullPath + "'") |
|
||||
} |
|
||||
n.handlers = handlers |
|
||||
n.fullPath = fullPath |
|
||||
return |
|
||||
} |
|
||||
} |
|
||||
|
|
||||
func (n *node) insertChild(path string, fullPath string, handlers HandlersChain) { |
|
||||
for { |
|
||||
// Find prefix until first wildcard
|
|
||||
wildcard, i, valid := findWildcard(path) |
|
||||
if i < 0 { // No wildcard found
|
|
||||
break |
|
||||
} |
|
||||
|
|
||||
// The wildcard name must only contain one ':' or '*' character
|
|
||||
if !valid { |
|
||||
panic("only one wildcard per path segment is allowed, has: '" + |
|
||||
wildcard + "' in path '" + fullPath + "'") |
|
||||
} |
|
||||
|
|
||||
// check if the wildcard has a name
|
|
||||
if len(wildcard) < 2 { |
|
||||
panic("wildcards must be named with a non-empty name in path '" + fullPath + "'") |
|
||||
} |
|
||||
|
|
||||
if wildcard[0] == ':' { // param
|
|
||||
if i > 0 { |
|
||||
// Insert prefix before the current wildcard
|
|
||||
n.path = path[:i] |
|
||||
path = path[i:] |
|
||||
} |
|
||||
|
|
||||
child := &node{ |
|
||||
nType: param, |
|
||||
path: wildcard, |
|
||||
fullPath: fullPath, |
|
||||
} |
|
||||
n.addChild(child) |
|
||||
n.wildChild = true |
|
||||
n = child |
|
||||
n.priority++ |
|
||||
|
|
||||
// if the path doesn't end with the wildcard, then there
|
|
||||
// will be another subpath starting with '/'
|
|
||||
if len(wildcard) < len(path) { |
|
||||
path = path[len(wildcard):] |
|
||||
|
|
||||
child := &node{ |
|
||||
priority: 1, |
|
||||
fullPath: fullPath, |
|
||||
} |
|
||||
n.addChild(child) |
|
||||
n = child |
|
||||
continue |
|
||||
} |
|
||||
|
|
||||
// Otherwise we're done. Insert the handle in the new leaf
|
|
||||
n.handlers = handlers |
|
||||
return |
|
||||
} |
|
||||
|
|
||||
// catchAll
|
|
||||
if i+len(wildcard) != len(path) { |
|
||||
panic("catch-all routes are only allowed at the end of the path in path '" + fullPath + "'") |
|
||||
} |
|
||||
|
|
||||
if len(n.path) > 0 && n.path[len(n.path)-1] == '/' { |
|
||||
pathSeg := strings.SplitN(n.children[0].path, "/", 2)[0] |
|
||||
panic("catch-all wildcard '" + path + |
|
||||
"' in new path '" + fullPath + |
|
||||
"' conflicts with existing path segment '" + pathSeg + |
|
||||
"' in existing prefix '" + n.path + pathSeg + |
|
||||
"'") |
|
||||
} |
|
||||
|
|
||||
// currently fixed width 1 for '/'
|
|
||||
i-- |
|
||||
if path[i] != '/' { |
|
||||
panic("no / before catch-all in path '" + fullPath + "'") |
|
||||
} |
|
||||
|
|
||||
n.path = path[:i] |
|
||||
|
|
||||
// First node: catchAll node with empty path
|
|
||||
child := &node{ |
|
||||
wildChild: true, |
|
||||
nType: catchAll, |
|
||||
fullPath: fullPath, |
|
||||
} |
|
||||
|
|
||||
n.addChild(child) |
|
||||
n.indices = string('/') |
|
||||
n = child |
|
||||
n.priority++ |
|
||||
|
|
||||
// second node: node holding the variable
|
|
||||
child = &node{ |
|
||||
path: path[i:], |
|
||||
nType: catchAll, |
|
||||
handlers: handlers, |
|
||||
priority: 1, |
|
||||
fullPath: fullPath, |
|
||||
} |
|
||||
n.children = []*node{child} |
|
||||
|
|
||||
return |
|
||||
} |
|
||||
|
|
||||
// If no wildcard was found, simply insert the path and handle
|
|
||||
n.path = path |
|
||||
n.handlers = handlers |
|
||||
n.fullPath = fullPath |
|
||||
} |
|
||||
func (n *node) incrementChildPrio(pos int) int { |
|
||||
cs := n.children |
|
||||
cs[pos].priority++ |
|
||||
prio := cs[pos].priority |
|
||||
|
|
||||
// Adjust position (move to front)
|
|
||||
newPos := pos |
|
||||
for ; newPos > 0 && cs[newPos-1].priority < prio; newPos-- { |
|
||||
// Swap node positions
|
|
||||
cs[newPos-1], cs[newPos] = cs[newPos], cs[newPos-1] |
|
||||
} |
|
||||
|
|
||||
// Build new index char string
|
|
||||
if newPos != pos { |
|
||||
n.indices = n.indices[:newPos] + // Unchanged prefix, might be empty
|
|
||||
n.indices[pos:pos+1] + // The index char we move
|
|
||||
n.indices[newPos:pos] + n.indices[pos+1:] // Rest without char at 'pos'
|
|
||||
} |
|
||||
|
|
||||
return newPos |
|
||||
} |
|
||||
|
|
||||
func (n *node) getValue(path string, params *Params, skippedNodes *[]skippedNode, unescape bool) (value nodeValue) { |
|
||||
var globalParamsCount int16 |
|
||||
|
|
||||
walk: // Outer loop for walking the tree
|
|
||||
for { |
|
||||
prefix := n.path |
|
||||
if len(path) > len(prefix) { |
|
||||
if path[:len(prefix)] == prefix { |
|
||||
path = path[len(prefix):] |
|
||||
|
|
||||
// Try all the non-wildcard children first by matching the indices
|
|
||||
idxc := path[0] |
|
||||
for i, c := range []byte(n.indices) { |
|
||||
if c == idxc { |
|
||||
// strings.HasPrefix(n.children[len(n.children)-1].path, ":") == n.wildChild
|
|
||||
if n.wildChild { |
|
||||
index := len(*skippedNodes) |
|
||||
*skippedNodes = (*skippedNodes)[:index+1] |
|
||||
(*skippedNodes)[index] = skippedNode{ |
|
||||
path: prefix + path, |
|
||||
node: &node{ |
|
||||
path: n.path, |
|
||||
wildChild: n.wildChild, |
|
||||
nType: n.nType, |
|
||||
priority: n.priority, |
|
||||
children: n.children, |
|
||||
handlers: n.handlers, |
|
||||
fullPath: n.fullPath, |
|
||||
}, |
|
||||
paramsCount: globalParamsCount, |
|
||||
} |
|
||||
} |
|
||||
|
|
||||
n = n.children[i] |
|
||||
continue walk |
|
||||
} |
|
||||
} |
|
||||
|
|
||||
if !n.wildChild { |
|
||||
// If the path at the end of the loop is not equal to '/' and the current node has no child nodes
|
|
||||
// the current node needs to roll back to last vaild skippedNode
|
|
||||
if path != "/" { |
|
||||
for l := len(*skippedNodes); l > 0; { |
|
||||
skippedNode := (*skippedNodes)[l-1] |
|
||||
*skippedNodes = (*skippedNodes)[:l-1] |
|
||||
if strings.HasSuffix(skippedNode.path, path) { |
|
||||
path = skippedNode.path |
|
||||
n = skippedNode.node |
|
||||
if value.params != nil { |
|
||||
*value.params = (*value.params)[:skippedNode.paramsCount] |
|
||||
} |
|
||||
globalParamsCount = skippedNode.paramsCount |
|
||||
continue walk |
|
||||
} |
|
||||
} |
|
||||
} |
|
||||
|
|
||||
// Nothing found.
|
|
||||
// We can recommend to redirect to the same URL without a
|
|
||||
// trailing slash if a leaf exists for that path.
|
|
||||
value.tsr = path == "/" && n.handlers != nil |
|
||||
return |
|
||||
} |
|
||||
|
|
||||
// Handle wildcard child, which is always at the end of the array
|
|
||||
n = n.children[len(n.children)-1] |
|
||||
globalParamsCount++ |
|
||||
|
|
||||
switch n.nType { |
|
||||
case param: |
|
||||
// fix truncate the parameter
|
|
||||
// tree_test.go line: 204
|
|
||||
|
|
||||
// Find param end (either '/' or path end)
|
|
||||
end := 0 |
|
||||
for end < len(path) && path[end] != '/' { |
|
||||
end++ |
|
||||
} |
|
||||
|
|
||||
// Save param value
|
|
||||
if params != nil && cap(*params) > 0 { |
|
||||
if value.params == nil { |
|
||||
value.params = params |
|
||||
} |
|
||||
// Expand slice within preallocated capacity
|
|
||||
i := len(*value.params) |
|
||||
*value.params = (*value.params)[:i+1] |
|
||||
val := path[:end] |
|
||||
if unescape { |
|
||||
if v, err := url.QueryUnescape(val); err == nil { |
|
||||
val = v |
|
||||
} |
|
||||
} |
|
||||
(*value.params)[i] = Param{ |
|
||||
Key: n.path[1:], |
|
||||
Value: val, |
|
||||
} |
|
||||
} |
|
||||
|
|
||||
// we need to go deeper!
|
|
||||
if end < len(path) { |
|
||||
if len(n.children) > 0 { |
|
||||
path = path[end:] |
|
||||
n = n.children[0] |
|
||||
continue walk |
|
||||
} |
|
||||
|
|
||||
// ... but we can't
|
|
||||
value.tsr = len(path) == end+1 |
|
||||
return |
|
||||
} |
|
||||
|
|
||||
if value.handlers = n.handlers; value.handlers != nil { |
|
||||
value.fullPath = n.fullPath |
|
||||
return |
|
||||
} |
|
||||
if len(n.children) == 1 { |
|
||||
// No handle found. Check if a handle for this path + a
|
|
||||
// trailing slash exists for TSR recommendation
|
|
||||
n = n.children[0] |
|
||||
value.tsr = (n.path == "/" && n.handlers != nil) || (n.path == "" && n.indices == "/") |
|
||||
} |
|
||||
return |
|
||||
|
|
||||
case catchAll: |
|
||||
// Save param value
|
|
||||
if params != nil { |
|
||||
if value.params == nil { |
|
||||
value.params = params |
|
||||
} |
|
||||
// Expand slice within preallocated capacity
|
|
||||
i := len(*value.params) |
|
||||
*value.params = (*value.params)[:i+1] |
|
||||
val := path |
|
||||
if unescape { |
|
||||
if v, err := url.QueryUnescape(path); err == nil { |
|
||||
val = v |
|
||||
} |
|
||||
} |
|
||||
(*value.params)[i] = Param{ |
|
||||
Key: n.path[2:], |
|
||||
Value: val, |
|
||||
} |
|
||||
} |
|
||||
|
|
||||
value.handlers = n.handlers |
|
||||
value.fullPath = n.fullPath |
|
||||
return |
|
||||
|
|
||||
default: |
|
||||
panic("invalid node type") |
|
||||
} |
|
||||
} |
|
||||
} |
|
||||
|
|
||||
if path == prefix { |
|
||||
// If the current path does not equal '/' and the node does not have a registered handle and the most recently matched node has a child node
|
|
||||
// the current node needs to roll back to last vaild skippedNode
|
|
||||
if n.handlers == nil && path != "/" { |
|
||||
for l := len(*skippedNodes); l > 0; { |
|
||||
skippedNode := (*skippedNodes)[l-1] |
|
||||
*skippedNodes = (*skippedNodes)[:l-1] |
|
||||
if strings.HasSuffix(skippedNode.path, path) { |
|
||||
path = skippedNode.path |
|
||||
n = skippedNode.node |
|
||||
if value.params != nil { |
|
||||
*value.params = (*value.params)[:skippedNode.paramsCount] |
|
||||
} |
|
||||
globalParamsCount = skippedNode.paramsCount |
|
||||
continue walk |
|
||||
} |
|
||||
} |
|
||||
// n = latestNode.children[len(latestNode.children)-1]
|
|
||||
} |
|
||||
// We should have reached the node containing the handle.
|
|
||||
// Check if this node has a handle registered.
|
|
||||
if value.handlers = n.handlers; value.handlers != nil { |
|
||||
value.fullPath = n.fullPath |
|
||||
return |
|
||||
} |
|
||||
|
|
||||
// If there is no handle for this route, but this route has a
|
|
||||
// wildcard child, there must be a handle for this path with an
|
|
||||
// additional trailing slash
|
|
||||
if path == "/" && n.wildChild && n.nType != root { |
|
||||
value.tsr = true |
|
||||
return |
|
||||
} |
|
||||
|
|
||||
// No handle found. Check if a handle for this path + a
|
|
||||
// trailing slash exists for trailing slash recommendation
|
|
||||
for i, c := range []byte(n.indices) { |
|
||||
if c == '/' { |
|
||||
n = n.children[i] |
|
||||
value.tsr = (len(n.path) == 1 && n.handlers != nil) || |
|
||||
(n.nType == catchAll && n.children[0].handlers != nil) |
|
||||
return |
|
||||
} |
|
||||
} |
|
||||
|
|
||||
return |
|
||||
} |
|
||||
|
|
||||
// Nothing found. We can recommend to redirect to the same URL with an
|
|
||||
// extra trailing slash if a leaf exists for that path
|
|
||||
value.tsr = path == "/" || |
|
||||
(len(prefix) == len(path)+1 && prefix[len(path)] == '/' && |
|
||||
path == prefix[:len(prefix)-1] && n.handlers != nil) |
|
||||
|
|
||||
// roll back to last valid skippedNode
|
|
||||
if !value.tsr && path != "/" { |
|
||||
for l := len(*skippedNodes); l > 0; { |
|
||||
skippedNode := (*skippedNodes)[l-1] |
|
||||
*skippedNodes = (*skippedNodes)[:l-1] |
|
||||
if strings.HasSuffix(skippedNode.path, path) { |
|
||||
path = skippedNode.path |
|
||||
n = skippedNode.node |
|
||||
if value.params != nil { |
|
||||
*value.params = (*value.params)[:skippedNode.paramsCount] |
|
||||
} |
|
||||
globalParamsCount = skippedNode.paramsCount |
|
||||
continue walk |
|
||||
} |
|
||||
} |
|
||||
} |
|
||||
|
|
||||
return |
|
||||
} |
|
||||
} |
|
||||
|
|
||||
func (n *node) findCaseInsensitivePath(path string, fixTrailingSlash bool) ([]byte, bool) { |
|
||||
const stackBufSize = 128 |
|
||||
|
|
||||
// Use a static sized buffer on the stack in the common case.
|
|
||||
// If the path is too long, allocate a buffer on the heap instead.
|
|
||||
buf := make([]byte, 0, stackBufSize) |
|
||||
if length := len(path) + 1; length > stackBufSize { |
|
||||
buf = make([]byte, 0, length) |
|
||||
} |
|
||||
|
|
||||
ciPath := n.findCaseInsensitivePathRec( |
|
||||
path, |
|
||||
buf, // Preallocate enough memory for new path
|
|
||||
[4]byte{}, // Empty rune buffer
|
|
||||
fixTrailingSlash, |
|
||||
) |
|
||||
|
|
||||
return ciPath, ciPath != nil |
|
||||
} |
|
||||
|
|
||||
// 使用的递归不区分大小写查找函数
|
|
||||
func (n *node) findCaseInsensitivePathRec(path string, ciPath []byte, rb [4]byte, fixTrailingSlash bool) []byte { |
|
||||
npLen := len(n.path) |
|
||||
|
|
||||
walk: // Outer loop for walking the tree
|
|
||||
for len(path) >= npLen && (npLen == 0 || strings.EqualFold(path[1:npLen], n.path[1:])) { |
|
||||
// Add common prefix to result
|
|
||||
oldPath := path |
|
||||
path = path[npLen:] |
|
||||
ciPath = append(ciPath, n.path...) |
|
||||
|
|
||||
if len(path) == 0 { |
|
||||
// We should have reached the node containing the handle.
|
|
||||
// Check if this node has a handle registered.
|
|
||||
if n.handlers != nil { |
|
||||
return ciPath |
|
||||
} |
|
||||
|
|
||||
// No handle found.
|
|
||||
// Try to fix the path by adding a trailing slash
|
|
||||
if fixTrailingSlash { |
|
||||
for i, c := range []byte(n.indices) { |
|
||||
if c == '/' { |
|
||||
n = n.children[i] |
|
||||
if (len(n.path) == 1 && n.handlers != nil) || |
|
||||
(n.nType == catchAll && n.children[0].handlers != nil) { |
|
||||
return append(ciPath, '/') |
|
||||
} |
|
||||
return nil |
|
||||
} |
|
||||
} |
|
||||
} |
|
||||
return nil |
|
||||
} |
|
||||
|
|
||||
// If this node does not have a wildcard (param or catchAll) child,
|
|
||||
// we can just look up the next child node and continue to walk down
|
|
||||
// the tree
|
|
||||
if !n.wildChild { |
|
||||
// Skip rune bytes already processed
|
|
||||
rb = shiftNRuneBytes(rb, npLen) |
|
||||
|
|
||||
if rb[0] != 0 { |
|
||||
// Old rune not finished
|
|
||||
idxc := rb[0] |
|
||||
for i, c := range []byte(n.indices) { |
|
||||
if c == idxc { |
|
||||
// continue with child node
|
|
||||
n = n.children[i] |
|
||||
npLen = len(n.path) |
|
||||
continue walk |
|
||||
} |
|
||||
} |
|
||||
} else { |
|
||||
// Process a new rune
|
|
||||
var rv rune |
|
||||
|
|
||||
// Find rune start.
|
|
||||
// Runes are up to 4 byte long,
|
|
||||
// -4 would definitely be another rune.
|
|
||||
var off int |
|
||||
for max := min(npLen, 3); off < max; off++ { |
|
||||
if i := npLen - off; utf8.RuneStart(oldPath[i]) { |
|
||||
// read rune from cached path
|
|
||||
rv, _ = utf8.DecodeRuneInString(oldPath[i:]) |
|
||||
break |
|
||||
} |
|
||||
} |
|
||||
|
|
||||
// Calculate lowercase bytes of current rune
|
|
||||
lo := unicode.ToLower(rv) |
|
||||
utf8.EncodeRune(rb[:], lo) |
|
||||
|
|
||||
// Skip already processed bytes
|
|
||||
rb = shiftNRuneBytes(rb, off) |
|
||||
|
|
||||
idxc := rb[0] |
|
||||
for i, c := range []byte(n.indices) { |
|
||||
// Lowercase matches
|
|
||||
if c == idxc { |
|
||||
// must use a recursive approach since both the
|
|
||||
// uppercase byte and the lowercase byte might exist
|
|
||||
// as an index
|
|
||||
if out := n.children[i].findCaseInsensitivePathRec( |
|
||||
path, ciPath, rb, fixTrailingSlash, |
|
||||
); out != nil { |
|
||||
return out |
|
||||
} |
|
||||
break |
|
||||
} |
|
||||
} |
|
||||
|
|
||||
// If we found no match, the same for the uppercase rune,
|
|
||||
// if it differs
|
|
||||
if up := unicode.ToUpper(rv); up != lo { |
|
||||
utf8.EncodeRune(rb[:], up) |
|
||||
rb = shiftNRuneBytes(rb, off) |
|
||||
|
|
||||
idxc := rb[0] |
|
||||
for i, c := range []byte(n.indices) { |
|
||||
// Uppercase matches
|
|
||||
if c == idxc { |
|
||||
// Continue with child node
|
|
||||
n = n.children[i] |
|
||||
npLen = len(n.path) |
|
||||
continue walk |
|
||||
} |
|
||||
} |
|
||||
} |
|
||||
} |
|
||||
|
|
||||
// Nothing found. We can recommend to redirect to the same URL
|
|
||||
// without a trailing slash if a leaf exists for that path
|
|
||||
if fixTrailingSlash && path == "/" && n.handlers != nil { |
|
||||
return ciPath |
|
||||
} |
|
||||
return nil |
|
||||
} |
|
||||
|
|
||||
n = n.children[0] |
|
||||
switch n.nType { |
|
||||
case param: |
|
||||
// Find param end (either '/' or path end)
|
|
||||
end := 0 |
|
||||
for end < len(path) && path[end] != '/' { |
|
||||
end++ |
|
||||
} |
|
||||
|
|
||||
// Add param value to case insensitive path
|
|
||||
ciPath = append(ciPath, path[:end]...) |
|
||||
|
|
||||
// We need to go deeper!
|
|
||||
if end < len(path) { |
|
||||
if len(n.children) > 0 { |
|
||||
// Continue with child node
|
|
||||
n = n.children[0] |
|
||||
npLen = len(n.path) |
|
||||
path = path[end:] |
|
||||
continue |
|
||||
} |
|
||||
|
|
||||
// ... but we can't
|
|
||||
if fixTrailingSlash && len(path) == end+1 { |
|
||||
return ciPath |
|
||||
} |
|
||||
return nil |
|
||||
} |
|
||||
|
|
||||
if n.handlers != nil { |
|
||||
return ciPath |
|
||||
} |
|
||||
|
|
||||
if fixTrailingSlash && len(n.children) == 1 { |
|
||||
// No handle found. Check if a handle for this path + a
|
|
||||
// trailing slash exists
|
|
||||
n = n.children[0] |
|
||||
if n.path == "/" && n.handlers != nil { |
|
||||
return append(ciPath, '/') |
|
||||
} |
|
||||
} |
|
||||
|
|
||||
return nil |
|
||||
|
|
||||
case catchAll: |
|
||||
return append(ciPath, path...) |
|
||||
|
|
||||
default: |
|
||||
panic("invalid node type") |
|
||||
} |
|
||||
} |
|
||||
|
|
||||
// Nothing found.
|
|
||||
// Try to fix the path by adding / removing a trailing slash
|
|
||||
if fixTrailingSlash { |
|
||||
if path == "/" { |
|
||||
return ciPath |
|
||||
} |
|
||||
if len(path)+1 == npLen && n.path[len(path)] == '/' && |
|
||||
strings.EqualFold(path[1:], n.path[1:len(path)]) && n.handlers != nil { |
|
||||
return append(ciPath, n.path...) |
|
||||
} |
|
||||
} |
|
||||
return nil |
|
||||
} |
|
||||
|
|
||||
type methodTree struct { |
|
||||
method string |
|
||||
root *node |
|
||||
} |
|
||||
|
|
||||
type methodTrees []methodTree |
|
||||
|
|
||||
func (trees methodTrees) get(method string) *node { |
|
||||
for _, tree := range trees { |
|
||||
if tree.method == method { |
|
||||
return tree.root |
|
||||
} |
|
||||
} |
|
||||
return nil |
|
||||
} |
|
||||
|
|
||||
/* |
|
||||
将数组中的字节向左移动 n 个字节 |
|
||||
*/ |
|
||||
func shiftNRuneBytes(rb [4]byte, n int) [4]byte { |
|
||||
switch n { |
|
||||
case 0: |
|
||||
return rb |
|
||||
case 1: |
|
||||
return [4]byte{rb[1], rb[2], rb[3], 0} |
|
||||
case 2: |
|
||||
return [4]byte{rb[2], rb[3]} |
|
||||
case 3: |
|
||||
return [4]byte{rb[3]} |
|
||||
default: |
|
||||
return [4]byte{} |
|
||||
} |
|
||||
} |
|
||||
|
|
||||
func longestCommonPrefix(a, b string) int { |
|
||||
i := 0 |
|
||||
max := min(len(a), len(b)) |
|
||||
for i < max && a[i] == b[i] { |
|
||||
i++ |
|
||||
} |
|
||||
return i |
|
||||
} |
|
||||
|
|
||||
func findWildcard(path string) (wildcard string, i int, valid bool) { |
|
||||
// Find start
|
|
||||
for start, c := range []byte(path) { |
|
||||
// A wildcard starts with ':' (param) or '*' (catch-all)
|
|
||||
if c != ':' && c != '*' { |
|
||||
continue |
|
||||
} |
|
||||
|
|
||||
// Find end and check for invalid characters
|
|
||||
valid = true |
|
||||
for end, c := range []byte(path[start+1:]) { |
|
||||
switch c { |
|
||||
case '/': |
|
||||
return path[start : start+1+end], start, valid |
|
||||
case ':', '*': |
|
||||
valid = false |
|
||||
} |
|
||||
} |
|
||||
return path[start:], start, valid |
|
||||
} |
|
||||
return "", -1, false |
|
||||
} |
|
||||
|
|
||||
func countParams(path string) uint16 { |
|
||||
var n uint16 |
|
||||
s := utils.StringToBytes(path) |
|
||||
n += uint16(bytes.Count(s, strColon)) |
|
||||
n += uint16(bytes.Count(s, strStar)) |
|
||||
return n |
|
||||
} |
|
||||
|
|
||||
func countSections(path string) uint16 { |
|
||||
s := utils.StringToBytes(path) |
|
||||
return uint16(bytes.Count(s, strSlash)) |
|
||||
} |
|
||||
func min(a, b int) int { |
|
||||
if a <= b { |
|
||||
return a |
|
||||
} |
|
||||
return b |
|
||||
} |
|
||||
@ -1,201 +0,0 @@ |
|||||
package engine |
|
||||
|
|
||||
import ( |
|
||||
"net" |
|
||||
"path" |
|
||||
"reflect" |
|
||||
"runtime" |
|
||||
"strings" |
|
||||
"unicode" |
|
||||
) |
|
||||
|
|
||||
/* |
|
||||
解析一个 IP 的字符串表示并返回一个 net.IP |
|
||||
最小字节表示,如果输入无效,则为零。 |
|
||||
*/ |
|
||||
func parseIP(ip string) net.IP { |
|
||||
parsedIP := net.ParseIP(ip) |
|
||||
if ipv4 := parsedIP.To4(); ipv4 != nil { |
|
||||
return ipv4 |
|
||||
} |
|
||||
return parsedIP |
|
||||
} |
|
||||
|
|
||||
func lastChar(str string) uint8 { |
|
||||
if str == "" { |
|
||||
panic("The length of the string can't be 0") |
|
||||
} |
|
||||
return str[len(str)-1] |
|
||||
} |
|
||||
|
|
||||
func assert1(guard bool, text string) { |
|
||||
if !guard { |
|
||||
panic(text) |
|
||||
} |
|
||||
} |
|
||||
|
|
||||
func nameOfFunction(f interface{}) string { |
|
||||
return runtime.FuncForPC(reflect.ValueOf(f).Pointer()).Name() |
|
||||
} |
|
||||
|
|
||||
func joinPaths(absolutePath, relativePath string) string { |
|
||||
if relativePath == "" { |
|
||||
return absolutePath |
|
||||
} |
|
||||
|
|
||||
finalPath := path.Join(absolutePath, relativePath) |
|
||||
if lastChar(relativePath) == '/' && lastChar(finalPath) != '/' { |
|
||||
return finalPath + "/" |
|
||||
} |
|
||||
return finalPath |
|
||||
} |
|
||||
|
|
||||
func iterate(path, method string, routes RoutesInfo, root *node) RoutesInfo { |
|
||||
path += root.path |
|
||||
if len(root.handlers) > 0 { |
|
||||
handlerFunc := root.handlers.Last() |
|
||||
routes = append(routes, RouteInfo{ |
|
||||
Method: method, |
|
||||
Path: path, |
|
||||
Handler: nameOfFunction(handlerFunc), |
|
||||
HandlerFunc: handlerFunc, |
|
||||
}) |
|
||||
} |
|
||||
for _, child := range root.children { |
|
||||
routes = iterate(path, method, routes, child) |
|
||||
} |
|
||||
return routes |
|
||||
} |
|
||||
func filterFlags(content string) string { |
|
||||
for i, char := range content { |
|
||||
if char == ' ' || char == ';' { |
|
||||
return content[:i] |
|
||||
} |
|
||||
} |
|
||||
return content |
|
||||
} |
|
||||
|
|
||||
func cleanPath(p string) string { |
|
||||
const stackBufSize = 128 |
|
||||
if p == "" { |
|
||||
return "/" |
|
||||
} |
|
||||
buf := make([]byte, 0, stackBufSize) |
|
||||
n := len(p) |
|
||||
r := 1 |
|
||||
w := 1 |
|
||||
if p[0] != '/' { |
|
||||
r = 0 |
|
||||
if n+1 > stackBufSize { |
|
||||
buf = make([]byte, n+1) |
|
||||
} else { |
|
||||
buf = buf[:n+1] |
|
||||
} |
|
||||
buf[0] = '/' |
|
||||
} |
|
||||
|
|
||||
trailing := n > 1 && p[n-1] == '/' |
|
||||
|
|
||||
for r < n { |
|
||||
switch { |
|
||||
case p[r] == '/': |
|
||||
r++ |
|
||||
case p[r] == '.' && r+1 == n: |
|
||||
trailing = true |
|
||||
r++ |
|
||||
case p[r] == '.' && p[r+1] == '/': |
|
||||
r += 2 |
|
||||
|
|
||||
case p[r] == '.' && p[r+1] == '.' && (r+2 == n || p[r+2] == '/'): |
|
||||
r += 3 |
|
||||
if w > 1 { |
|
||||
w-- |
|
||||
if len(buf) == 0 { |
|
||||
for w > 1 && p[w] != '/' { |
|
||||
w-- |
|
||||
} |
|
||||
} else { |
|
||||
for w > 1 && buf[w] != '/' { |
|
||||
w-- |
|
||||
} |
|
||||
} |
|
||||
} |
|
||||
default: |
|
||||
if w > 1 { |
|
||||
bufApp(&buf, p, w, '/') |
|
||||
w++ |
|
||||
} |
|
||||
for r < n && p[r] != '/' { |
|
||||
bufApp(&buf, p, w, p[r]) |
|
||||
w++ |
|
||||
r++ |
|
||||
} |
|
||||
} |
|
||||
} |
|
||||
if trailing && w > 1 { |
|
||||
bufApp(&buf, p, w, '/') |
|
||||
w++ |
|
||||
} |
|
||||
if len(buf) == 0 { |
|
||||
return p[:w] |
|
||||
} |
|
||||
return string(buf[:w]) |
|
||||
} |
|
||||
|
|
||||
func bufApp(buf *[]byte, s string, w int, c byte) { |
|
||||
b := *buf |
|
||||
if len(b) == 0 { |
|
||||
// No modification of the original string so far.
|
|
||||
// If the next character is the same as in the original string, we do
|
|
||||
// not yet have to allocate a buffer.
|
|
||||
if s[w] == c { |
|
||||
return |
|
||||
} |
|
||||
|
|
||||
// Otherwise use either the stack buffer, if it is large enough, or
|
|
||||
// allocate a new buffer on the heap, and copy all previous characters.
|
|
||||
length := len(s) |
|
||||
if length > cap(b) { |
|
||||
*buf = make([]byte, length) |
|
||||
} else { |
|
||||
*buf = (*buf)[:length] |
|
||||
} |
|
||||
b = *buf |
|
||||
|
|
||||
copy(b, s[:w]) |
|
||||
} |
|
||||
b[w] = c |
|
||||
} |
|
||||
|
|
||||
func parseAccept(acceptHeader string) []string { |
|
||||
parts := strings.Split(acceptHeader, ",") |
|
||||
out := make([]string, 0, len(parts)) |
|
||||
for _, part := range parts { |
|
||||
if i := strings.IndexByte(part, ';'); i > 0 { |
|
||||
part = part[:i] |
|
||||
} |
|
||||
if part = strings.TrimSpace(part); part != "" { |
|
||||
out = append(out, part) |
|
||||
} |
|
||||
} |
|
||||
return out |
|
||||
} |
|
||||
|
|
||||
func chooseData(custom, wildcard interface{}) interface{} { |
|
||||
if custom != nil { |
|
||||
return custom |
|
||||
} |
|
||||
if wildcard != nil { |
|
||||
return wildcard |
|
||||
} |
|
||||
panic("negotiation config is invalid") |
|
||||
} |
|
||||
|
|
||||
func isASCII(s string) bool { |
|
||||
for i := 0; i < len(s); i++ { |
|
||||
if s[i] > unicode.MaxASCII { |
|
||||
return false |
|
||||
} |
|
||||
} |
|
||||
return true |
|
||||
} |
|
||||
@ -1,171 +0,0 @@ |
|||||
package gin |
|
||||
|
|
||||
import ( |
|
||||
"context" |
|
||||
"errors" |
|
||||
"fmt" |
|
||||
"net" |
|
||||
"net/http" |
|
||||
|
|
||||
"yunyan/lego/sys/gin/engine" |
|
||||
"yunyan/lego/sys/gin/middleware/logger" |
|
||||
"yunyan/lego/sys/gin/middleware/recovery" |
|
||||
|
|
||||
"github.com/gin-gonic/autotls" |
|
||||
) |
|
||||
|
|
||||
func newSys(options *Options) (sys *Gin, err error) { |
|
||||
sys = &Gin{ |
|
||||
options: options, |
|
||||
} |
|
||||
sys.engine = engine.NewEngine(engine.SetLog(options.Log), engine.SetMultipartMemory(options.MultipartMemory)) |
|
||||
///添加基础中间件
|
|
||||
sys.engine.Use(logger.Logger([]string{}), recovery.Recovery()) |
|
||||
if options.CertFile != "" && options.KeyFile != "" { |
|
||||
sys.RunTLS(options.ListenPort, options.CertFile, options.KeyFile) |
|
||||
} else if options.LetEncrypt { |
|
||||
sys.RunLetEncrypt(options.Domain...) |
|
||||
} else { |
|
||||
sys.Run(options.ListenPort) |
|
||||
} |
|
||||
return |
|
||||
} |
|
||||
|
|
||||
type Gin struct { |
|
||||
options *Options |
|
||||
server *http.Server |
|
||||
engine *engine.Engine |
|
||||
} |
|
||||
|
|
||||
func (this *Gin) Run(listenPort int) (err error) { |
|
||||
// if this.engine.IsUnsafeTrustedProxies() {
|
|
||||
// this.Warnf("You trusted all proxies, this is NOT safe. We recommend you to set a value.\n" +
|
|
||||
// "Please check https://pkg.go.dev/github.com/gin-gonic/gin#readme-don-t-trust-all-proxies for details.")
|
|
||||
|
|
||||
// }
|
|
||||
this.options.Log.Debugf("Listening and serving HTTP on:%d", listenPort) |
|
||||
this.server = &http.Server{ |
|
||||
Addr: fmt.Sprintf(":%d", listenPort), |
|
||||
Handler: this.engine.Handler(), |
|
||||
} |
|
||||
go func() { |
|
||||
if err := this.server.ListenAndServe(); err != nil && errors.Is(err, http.ErrServerClosed) { |
|
||||
this.options.Log.Errorln(err) |
|
||||
} |
|
||||
}() |
|
||||
// err = http.ListenAndServe(fmt.Sprintf(":%d", this.options.ListenPort), this.Handler())
|
|
||||
return |
|
||||
} |
|
||||
|
|
||||
func (this *Gin) RunTLS(listenPort int, certFile, keyFile string) (err error) { |
|
||||
this.options.Log.Debugf("Listening and serving HTTPS on :%d", listenPort) |
|
||||
// if this.engine.IsUnsafeTrustedProxies() {
|
|
||||
// this.Warnf("You trusted all proxies, this is NOT safe. We recommend you to set a value.\n" +
|
|
||||
// "Please check https://pkg.go.dev/github.com/gin-gonic/gin#readme-don-t-trust-all-proxies for details.")
|
|
||||
// }
|
|
||||
this.server = &http.Server{ |
|
||||
Addr: fmt.Sprintf(":%d", listenPort), |
|
||||
Handler: this.engine.Handler(), |
|
||||
} |
|
||||
go func() { |
|
||||
if err := this.server.ListenAndServeTLS(certFile, keyFile); err != nil && errors.Is(err, http.ErrServerClosed) { |
|
||||
this.options.Log.Errorln(err) |
|
||||
} |
|
||||
}() |
|
||||
// err = http.ListenAndServeTLS(addr, certFile, keyFile, this.Handler())
|
|
||||
return |
|
||||
} |
|
||||
|
|
||||
func (this *Gin) RunLetEncrypt(domain ...string) { |
|
||||
this.options.Log.Debugf("Listening and serving LetEncrypt on :%v", domain) |
|
||||
go func() { |
|
||||
if err := autotls.Run(this.engine, domain...); err != nil { |
|
||||
this.options.Log.Errorln(err) |
|
||||
} |
|
||||
}() |
|
||||
} |
|
||||
|
|
||||
func (this *Gin) RunListener(listener net.Listener) (err error) { |
|
||||
this.options.Log.Debugf("Listening and serving HTTP on listener what's bind with address@%s", listener.Addr()) |
|
||||
defer func() { |
|
||||
if err != nil { |
|
||||
this.options.Log.Errorln(err) |
|
||||
} |
|
||||
}() |
|
||||
// if this.engine.IsUnsafeTrustedProxies() {
|
|
||||
// this.Warnf("You trusted all proxies, this is NOT safe. We recommend you to set a value.\n" +
|
|
||||
// "Please check https://pkg.go.dev/github.com/gin-gonic/gin#readme-don-t-trust-all-proxies for details.")
|
|
||||
|
|
||||
// }
|
|
||||
err = http.Serve(listener, this.engine.Handler()) |
|
||||
return |
|
||||
} |
|
||||
func (this *Gin) HandleContext(c *engine.Context) { |
|
||||
this.HandleContext(c) |
|
||||
} |
|
||||
func (this *Gin) LoadHTMLGlob(pattern string) { |
|
||||
this.engine.LoadHTMLGlob(pattern) |
|
||||
} |
|
||||
|
|
||||
func (this *Gin) Close() (err error) { |
|
||||
if err = this.server.Shutdown(context.Background()); err != nil { |
|
||||
this.options.Log.Errorln(err) |
|
||||
} |
|
||||
this.server.Close() |
|
||||
return |
|
||||
} |
|
||||
func (this *Gin) Register(rcvr interface{}) { |
|
||||
this.engine.Register(rcvr) |
|
||||
} |
|
||||
|
|
||||
func (this *Gin) NoRoute(handlers ...engine.HandlerFunc) { |
|
||||
this.engine.NoRoute(handlers...) |
|
||||
} |
|
||||
|
|
||||
func (this *Gin) Group(relativePath string, handlers ...engine.HandlerFunc) engine.IRoutes { |
|
||||
return this.engine.Group(relativePath, handlers...) |
|
||||
} |
|
||||
func (this *Gin) Use(handlers ...engine.HandlerFunc) engine.IRoutes { |
|
||||
return this.engine.Use(handlers...) |
|
||||
} |
|
||||
func (this *Gin) Handle(httpMethod string, relativePath string, handlers ...engine.HandlerFunc) engine.IRoutes { |
|
||||
return this.engine.Handle(httpMethod, relativePath, handlers...) |
|
||||
} |
|
||||
func (this *Gin) Any(relativePath string, handlers ...engine.HandlerFunc) engine.IRoutes { |
|
||||
return this.engine.Any(relativePath, handlers...) |
|
||||
} |
|
||||
func (this *Gin) GET(httpMethod string, handlers ...engine.HandlerFunc) engine.IRoutes { |
|
||||
return this.engine.GET(httpMethod, handlers...) |
|
||||
} |
|
||||
func (this *Gin) POST(httpMethod string, handlers ...engine.HandlerFunc) engine.IRoutes { |
|
||||
return this.engine.POST(httpMethod, handlers...) |
|
||||
} |
|
||||
func (this *Gin) DELETE(httpMethod string, handlers ...engine.HandlerFunc) engine.IRoutes { |
|
||||
return this.engine.DELETE(httpMethod, handlers...) |
|
||||
} |
|
||||
func (this *Gin) PATCH(httpMethod string, handlers ...engine.HandlerFunc) engine.IRoutes { |
|
||||
return defsys.PATCH(httpMethod, handlers...) |
|
||||
} |
|
||||
func (this *Gin) PUT(httpMethod string, handlers ...engine.HandlerFunc) engine.IRoutes { |
|
||||
return this.engine.PUT(httpMethod, handlers...) |
|
||||
} |
|
||||
func (this *Gin) OPTIONS(httpMethod string, handlers ...engine.HandlerFunc) engine.IRoutes { |
|
||||
return this.engine.OPTIONS(httpMethod, handlers...) |
|
||||
} |
|
||||
func (this *Gin) HEAD(httpMethod string, handlers ...engine.HandlerFunc) engine.IRoutes { |
|
||||
return this.engine.HEAD(httpMethod, handlers...) |
|
||||
} |
|
||||
func (this *Gin) StaticFile(relativePath string, filepath string) engine.IRoutes { |
|
||||
return this.engine.StaticFile(relativePath, filepath) |
|
||||
} |
|
||||
func (this *Gin) StaticFileFS(relativePath string, filepath string, fs http.FileSystem) engine.IRoutes { |
|
||||
return this.engine.StaticFileFS(relativePath, filepath, fs) |
|
||||
} |
|
||||
|
|
||||
func (this *Gin) Static(relativePath string, root string) engine.IRoutes { |
|
||||
return this.engine.Static(relativePath, root) |
|
||||
} |
|
||||
|
|
||||
func (this *Gin) StaticFS(relativePath string, fs http.FileSystem) engine.IRoutes { |
|
||||
return this.engine.StaticFS(relativePath, fs) |
|
||||
} |
|
||||
@ -1,28 +0,0 @@ |
|||||
/* |
|
||||
解决跨域 中间件 |
|
||||
*/ |
|
||||
package cross |
|
||||
|
|
||||
import ( |
|
||||
"net/http" |
|
||||
|
|
||||
"yunyan/lego/sys/gin/engine" |
|
||||
) |
|
||||
|
|
||||
func handlerCors() engine.HandlerFunc { |
|
||||
return func(c *engine.Context) { |
|
||||
method := c.Request.Method |
|
||||
origin := c.Request.Header.Get("Origin") //请求头部
|
|
||||
if origin != "" { |
|
||||
c.Header("Access-Control-Allow-Origin", origin) |
|
||||
c.Header("Access-Control-Allow-Methods", "POST, GET, OPTIONS, PUT, DELETE, UPDATE") |
|
||||
c.Header("Access-Control-Allow-Headers", "Origin, X-Requested-With, Content-Type,X-Token, Accept, Authorization") |
|
||||
c.Header("Access-Control-Expose-Headers", "Content-Length, Access-Control-Allow-Origin, Access-Control-Allow-Headers, Cache-Control, Content-Language, Content-Type") |
|
||||
c.Header("Access-Control-Allow-Credentials", "true") |
|
||||
} //允许类型校验
|
|
||||
if method == "OPTIONS" { |
|
||||
c.AbortWithStatus(http.StatusNoContent) |
|
||||
} |
|
||||
c.Next() |
|
||||
} |
|
||||
} |
|
||||
@ -1,88 +0,0 @@ |
|||||
package jwt |
|
||||
|
|
||||
import ( |
|
||||
"net/http" |
|
||||
"strings" |
|
||||
"time" |
|
||||
|
|
||||
"yunyan/lego/core" |
|
||||
"yunyan/lego/sys/gin/engine" |
|
||||
|
|
||||
"github.com/golang-jwt/jwt" |
|
||||
) |
|
||||
|
|
||||
func NewJWT(key, tokenKey string) *JWT { |
|
||||
return &JWT{ |
|
||||
jwtkey: []byte(key), |
|
||||
tokenKey: tokenKey, |
|
||||
} |
|
||||
} |
|
||||
|
|
||||
type JWT struct { |
|
||||
jwtkey []byte |
|
||||
tokenKey string |
|
||||
} |
|
||||
|
|
||||
// CreateToken 生成token
|
|
||||
func CreateToken(key, Id string) (string, error) { |
|
||||
expireTime := time.Now().Add(2 * time.Hour) //过期时间
|
|
||||
nowTime := time.Now() //当前时间
|
|
||||
claims := jwt.StandardClaims{ |
|
||||
Id: Id, //用户Id
|
|
||||
ExpiresAt: expireTime.Unix(), //过期时间戳
|
|
||||
IssuedAt: nowTime.Unix(), //当前时间戳
|
|
||||
Issuer: "blogLeo", //颁发者签名
|
|
||||
Subject: "userToken", //签名主题
|
|
||||
|
|
||||
} |
|
||||
tokenStruct := jwt.NewWithClaims(jwt.SigningMethodHS256, claims) |
|
||||
return tokenStruct.SignedString([]byte(key)) |
|
||||
} |
|
||||
|
|
||||
// CheckToken 验证token
|
|
||||
func (this *JWT) CheckToken(token string) (*jwt.StandardClaims, bool) { |
|
||||
tokenObj, _ := jwt.ParseWithClaims(token, &jwt.StandardClaims{}, func(token *jwt.Token) (interface{}, error) { |
|
||||
return this.jwtkey, nil |
|
||||
}) |
|
||||
if key, _ := tokenObj.Claims.(*jwt.StandardClaims); tokenObj.Valid { |
|
||||
return key, true |
|
||||
} else { |
|
||||
return key, false |
|
||||
} |
|
||||
} |
|
||||
|
|
||||
// JwtMiddleware jwt中间件
|
|
||||
func (this *JWT) JwtMiddleware() engine.HandlerFunc { |
|
||||
return func(c *engine.Context) { |
|
||||
//从请求头中获取token
|
|
||||
tokenStr := c.Request.Header.Get(this.tokenKey) |
|
||||
//用户不存在
|
|
||||
if tokenStr == "" { |
|
||||
c.JSON(http.StatusOK, engine.H{"code": core.ErrorCode_NoLogin, "msg": "用户不存在"}) |
|
||||
c.Abort() //阻止执行
|
|
||||
return |
|
||||
} |
|
||||
//token格式错误
|
|
||||
tokenSlice := strings.Split(tokenStr, ".") |
|
||||
if len(tokenSlice) != 3 { |
|
||||
c.JSON(http.StatusOK, engine.H{"code": core.ErrorCode_NoLogin, "msg": "token格式错误"}) |
|
||||
c.Abort() //阻止执行
|
|
||||
return |
|
||||
} |
|
||||
//验证token
|
|
||||
tokenStruck, ok := this.CheckToken(tokenStr) |
|
||||
if !ok { |
|
||||
c.JSON(http.StatusOK, engine.H{"code": core.ErrorCode_NoLogin, "msg": "token不正确"}) |
|
||||
c.Abort() //阻止执行
|
|
||||
return |
|
||||
} |
|
||||
//token超时
|
|
||||
if time.Now().Unix() > tokenStruck.ExpiresAt { |
|
||||
c.JSON(http.StatusOK, engine.H{"code": core.ErrorCode_NoLogin, "msg": "token过期"}) |
|
||||
c.Abort() //阻止执行
|
|
||||
return |
|
||||
} |
|
||||
c.SetUserId(tokenStruck.Id) |
|
||||
c.Next() |
|
||||
} |
|
||||
} |
|
||||
@ -1,24 +0,0 @@ |
|||||
package jwt |
|
||||
|
|
||||
import ( |
|
||||
"fmt" |
|
||||
"testing" |
|
||||
|
|
||||
"github.com/golang-jwt/jwt" |
|
||||
) |
|
||||
|
|
||||
func Test_ParamSign(t *testing.T) { |
|
||||
if token, err := CreateToken("mysces", "liwei1dao"); err != nil { |
|
||||
fmt.Printf("err:%v", err) |
|
||||
return |
|
||||
} else { |
|
||||
tokenObj, _ := jwt.ParseWithClaims(token, &jwt.StandardClaims{}, func(token *jwt.Token) (interface{}, error) { |
|
||||
return []byte("mysces"), nil |
|
||||
}) |
|
||||
if key, ok := tokenObj.Claims.(*jwt.StandardClaims); ok && tokenObj.Valid { |
|
||||
fmt.Printf("key:%v", key) |
|
||||
} else { |
|
||||
fmt.Printf("void") |
|
||||
} |
|
||||
} |
|
||||
} |
|
||||
@ -1,85 +0,0 @@ |
|||||
package logger |
|
||||
|
|
||||
import ( |
|
||||
"fmt" |
|
||||
"net/http" |
|
||||
"time" |
|
||||
|
|
||||
"yunyan/lego/sys/gin/engine" |
|
||||
) |
|
||||
|
|
||||
type LogFormatterParams struct { |
|
||||
Request *http.Request |
|
||||
// TimeStamp shows the time after the server returns a response.
|
|
||||
TimeStamp time.Time |
|
||||
// StatusCode is HTTP response code.
|
|
||||
StatusCode int |
|
||||
// Latency is how much time the server cost to process a certain request.
|
|
||||
Latency time.Duration |
|
||||
// ClientIP equals Context's ClientIP method.
|
|
||||
ClientIP string |
|
||||
// Method is the HTTP method given to the request.
|
|
||||
Method string |
|
||||
// Path is a path the client requests.
|
|
||||
Path string |
|
||||
// ErrorMessage is set if error has occurred in processing the request.
|
|
||||
ErrorMessage string |
|
||||
// isTerm shows whether gin's output descriptor refers to a terminal.
|
|
||||
isTerm bool |
|
||||
// BodySize is the size of the Response Body
|
|
||||
BodySize int |
|
||||
// Keys are the keys set on the request's context.
|
|
||||
Keys map[string]interface{} |
|
||||
} |
|
||||
|
|
||||
func Logger(SkipPaths []string) engine.HandlerFunc { |
|
||||
var skip map[string]struct{} |
|
||||
if length := len(SkipPaths); length > 0 { |
|
||||
skip = make(map[string]struct{}, length) |
|
||||
|
|
||||
for _, path := range SkipPaths { |
|
||||
skip[path] = struct{}{} |
|
||||
} |
|
||||
} |
|
||||
return func(c *engine.Context) { |
|
||||
// Start timer
|
|
||||
start := time.Now() |
|
||||
path := c.Request.URL.Path |
|
||||
raw := c.Request.URL.RawQuery |
|
||||
|
|
||||
// Process request
|
|
||||
c.Next() |
|
||||
|
|
||||
// Log only when path is not being skipped
|
|
||||
if _, ok := skip[path]; !ok { |
|
||||
param := LogFormatterParams{ |
|
||||
Request: c.Request, |
|
||||
Keys: c.Keys, |
|
||||
} |
|
||||
// Stop timer
|
|
||||
param.TimeStamp = time.Now() |
|
||||
param.Latency = param.TimeStamp.Sub(start) |
|
||||
|
|
||||
param.ClientIP = c.ClientIP() |
|
||||
param.Method = c.Request.Method |
|
||||
param.StatusCode = c.Writer.Status() |
|
||||
param.ErrorMessage = c.Errors.ByType(engine.ErrorTypePrivate).String() |
|
||||
|
|
||||
param.BodySize = c.Writer.Size() |
|
||||
|
|
||||
if raw != "" { |
|
||||
path = path + "?" + raw |
|
||||
} |
|
||||
param.Path = path |
|
||||
c.Log.Debugf(fmt.Sprintf("[FORMATTER TEST] %v | %3d | %13v | %15s | %-7s %s\n%s", |
|
||||
param.TimeStamp.Format("2006/01/02 - 15:04:05"), |
|
||||
param.StatusCode, |
|
||||
param.Latency, |
|
||||
param.ClientIP, |
|
||||
param.Method, |
|
||||
param.Path, |
|
||||
param.ErrorMessage, |
|
||||
)) |
|
||||
} |
|
||||
} |
|
||||
} |
|
||||
@ -1,159 +0,0 @@ |
|||||
package recovery |
|
||||
|
|
||||
import ( |
|
||||
"bytes" |
|
||||
"errors" |
|
||||
"fmt" |
|
||||
"io/ioutil" |
|
||||
"net" |
|
||||
"net/http" |
|
||||
"net/http/httputil" |
|
||||
"os" |
|
||||
"runtime" |
|
||||
"strings" |
|
||||
"time" |
|
||||
|
|
||||
"yunyan/lego/sys/gin/engine" |
|
||||
) |
|
||||
|
|
||||
const ( |
|
||||
green = "\033[97;42m" |
|
||||
white = "\033[90;47m" |
|
||||
yellow = "\033[90;43m" |
|
||||
red = "\033[97;41m" |
|
||||
blue = "\033[97;44m" |
|
||||
magenta = "\033[97;45m" |
|
||||
cyan = "\033[97;46m" |
|
||||
reset = "\033[0m" |
|
||||
) |
|
||||
|
|
||||
var ( |
|
||||
dunno = []byte("???") |
|
||||
centerDot = []byte("·") |
|
||||
dot = []byte(".") |
|
||||
slash = []byte("/") |
|
||||
) |
|
||||
|
|
||||
// RecoveryFunc defines the function passable to CustomRecovery.
|
|
||||
type RecoveryFunc func(c *engine.Context, err interface{}) |
|
||||
|
|
||||
func Recovery() engine.HandlerFunc { |
|
||||
return CustomRecoveryWithWriter(defaultHandleRecovery) |
|
||||
} |
|
||||
|
|
||||
func defaultHandleRecovery(c *engine.Context, err interface{}) { |
|
||||
c.AbortWithStatus(http.StatusInternalServerError) |
|
||||
} |
|
||||
|
|
||||
// CustomRecoveryWithWriter returns a middleware for a given writer that recovers from any panics and calls the provided handle func to handle it.
|
|
||||
func CustomRecoveryWithWriter(handle RecoveryFunc) engine.HandlerFunc { |
|
||||
return func(c *engine.Context) { |
|
||||
defer func() { |
|
||||
if err := recover(); err != nil { |
|
||||
// Check for a broken connection, as it is not really a
|
|
||||
// condition that warrants a panic stack trace.
|
|
||||
var brokenPipe bool |
|
||||
if ne, ok := err.(*net.OpError); ok { |
|
||||
var se *os.SyscallError |
|
||||
if errors.As(ne, &se) { |
|
||||
if strings.Contains(strings.ToLower(se.Error()), "broken pipe") || strings.Contains(strings.ToLower(se.Error()), "connection reset by peer") { |
|
||||
brokenPipe = true |
|
||||
} |
|
||||
} |
|
||||
} |
|
||||
|
|
||||
stack := stack(3) |
|
||||
httpRequest, _ := httputil.DumpRequest(c.Request, false) |
|
||||
headers := strings.Split(string(httpRequest), "\r\n") |
|
||||
for idx, header := range headers { |
|
||||
current := strings.Split(header, ":") |
|
||||
if current[0] == "Authorization" { |
|
||||
headers[idx] = current[0] + ": *" |
|
||||
} |
|
||||
} |
|
||||
headersToStr := strings.Join(headers, "\r\n") |
|
||||
if brokenPipe { |
|
||||
c.Log.Errorf("%s\n%s%s", err, headersToStr, reset) |
|
||||
} else { |
|
||||
c.Log.Errorf("[Recovery] %s panic recovered:\n%s\n%s\n%s%s", |
|
||||
timeFormat(time.Now()), headersToStr, err, stack, reset) |
|
||||
} |
|
||||
|
|
||||
if brokenPipe { |
|
||||
// If the connection is dead, we can't write a status to it.
|
|
||||
c.Error(err.(error)) // nolint: errcheck
|
|
||||
c.Abort() |
|
||||
} else { |
|
||||
handle(c, err) |
|
||||
} |
|
||||
} |
|
||||
}() |
|
||||
c.Next() |
|
||||
} |
|
||||
} |
|
||||
|
|
||||
// stack returns a nicely formatted stack frame, skipping skip frames.
|
|
||||
func stack(skip int) []byte { |
|
||||
buf := new(bytes.Buffer) // the returned data
|
|
||||
// As we loop, we open files and read them. These variables record the currently
|
|
||||
// loaded file.
|
|
||||
var lines [][]byte |
|
||||
var lastFile string |
|
||||
for i := skip; ; i++ { // Skip the expected number of frames
|
|
||||
pc, file, line, ok := runtime.Caller(i) |
|
||||
if !ok { |
|
||||
break |
|
||||
} |
|
||||
// Print this much at least. If we can't find the source, it won't show.
|
|
||||
fmt.Fprintf(buf, "%s:%d (0x%x)\n", file, line, pc) |
|
||||
if file != lastFile { |
|
||||
data, err := ioutil.ReadFile(file) |
|
||||
if err != nil { |
|
||||
continue |
|
||||
} |
|
||||
lines = bytes.Split(data, []byte{'\n'}) |
|
||||
lastFile = file |
|
||||
} |
|
||||
fmt.Fprintf(buf, "\t%s: %s\n", function(pc), source(lines, line)) |
|
||||
} |
|
||||
return buf.Bytes() |
|
||||
} |
|
||||
|
|
||||
// source returns a space-trimmed slice of the n'th line.
|
|
||||
func source(lines [][]byte, n int) []byte { |
|
||||
n-- // in stack trace, lines are 1-indexed but our array is 0-indexed
|
|
||||
if n < 0 || n >= len(lines) { |
|
||||
return dunno |
|
||||
} |
|
||||
return bytes.TrimSpace(lines[n]) |
|
||||
} |
|
||||
|
|
||||
// function returns, if possible, the name of the function containing the PC.
|
|
||||
func function(pc uintptr) []byte { |
|
||||
fn := runtime.FuncForPC(pc) |
|
||||
if fn == nil { |
|
||||
return dunno |
|
||||
} |
|
||||
name := []byte(fn.Name()) |
|
||||
// The name includes the path name to the package, which is unnecessary
|
|
||||
// since the file name is already included. Plus, it has center dots.
|
|
||||
// That is, we see
|
|
||||
// runtime/debug.*T·ptrmethod
|
|
||||
// and want
|
|
||||
// *T.ptrmethod
|
|
||||
// Also the package path might contain dot (e.g. code.google.com/...),
|
|
||||
// so first eliminate the path prefix
|
|
||||
if lastSlash := bytes.LastIndex(name, slash); lastSlash >= 0 { |
|
||||
name = name[lastSlash+1:] |
|
||||
} |
|
||||
if period := bytes.Index(name, dot); period >= 0 { |
|
||||
name = name[period+1:] |
|
||||
} |
|
||||
name = bytes.Replace(name, centerDot, dot, -1) |
|
||||
return name |
|
||||
} |
|
||||
|
|
||||
// timeFormat returns a customized time string for logger.
|
|
||||
func timeFormat(t time.Time) string { |
|
||||
return t.Format("2006/01/02 - 15:04:05") |
|
||||
} |
|
||||
@ -1,101 +0,0 @@ |
|||||
package gin |
|
||||
|
|
||||
import ( |
|
||||
"yunyan/lego/sys/log" |
|
||||
"yunyan/lego/utils/mapstructure" |
|
||||
) |
|
||||
|
|
||||
type Option func(*Options) |
|
||||
type Options struct { |
|
||||
ListenPort int //监听端口
|
|
||||
CertFile string //tls文件
|
|
||||
KeyFile string //tls文件
|
|
||||
LetEncrypt bool //Let's Encrypt 证书模式
|
|
||||
Domain []string //域名
|
|
||||
MultipartMemory int64 //文件上传最大尺寸
|
|
||||
Debug bool //日志是否开启
|
|
||||
Log log.ILogger |
|
||||
} |
|
||||
|
|
||||
func SetListenPort(v int) Option { |
|
||||
return func(o *Options) { |
|
||||
o.ListenPort = v |
|
||||
} |
|
||||
} |
|
||||
|
|
||||
func SetCertFile(v string) Option { |
|
||||
return func(o *Options) { |
|
||||
o.CertFile = v |
|
||||
} |
|
||||
} |
|
||||
|
|
||||
func SetKeyFile(v string) Option { |
|
||||
return func(o *Options) { |
|
||||
o.KeyFile = v |
|
||||
} |
|
||||
} |
|
||||
|
|
||||
func SetLetEncrypt(v bool) Option { |
|
||||
return func(o *Options) { |
|
||||
o.LetEncrypt = v |
|
||||
} |
|
||||
} |
|
||||
|
|
||||
func SetDomain(v []string) Option { |
|
||||
return func(o *Options) { |
|
||||
o.Domain = v |
|
||||
} |
|
||||
} |
|
||||
|
|
||||
func SetMultipartMemory(v int64) Option { |
|
||||
return func(o *Options) { |
|
||||
o.MultipartMemory = v |
|
||||
} |
|
||||
} |
|
||||
|
|
||||
func SetDebug(v bool) Option { |
|
||||
return func(o *Options) { |
|
||||
o.Debug = v |
|
||||
} |
|
||||
} |
|
||||
|
|
||||
func SetLog(v log.ILogger) Option { |
|
||||
return func(o *Options) { |
|
||||
o.Log = v |
|
||||
} |
|
||||
} |
|
||||
|
|
||||
func newOptions(config map[string]interface{}, opts ...Option) (options *Options, err error) { |
|
||||
options = &Options{ |
|
||||
ListenPort: 8080, |
|
||||
CertFile: "", |
|
||||
KeyFile: "", |
|
||||
Debug: true, |
|
||||
} |
|
||||
if config != nil { |
|
||||
mapstructure.Decode(config, &options) |
|
||||
} |
|
||||
for _, o := range opts { |
|
||||
o(options) |
|
||||
} |
|
||||
if options.Log == nil { |
|
||||
options.Log = log.NewTurnlog(options.Debug, log.Clone("sys.gin", 3)) |
|
||||
} |
|
||||
return |
|
||||
} |
|
||||
|
|
||||
func newOptionsByOption(opts ...Option) (options *Options, err error) { |
|
||||
options = &Options{ |
|
||||
ListenPort: 8080, |
|
||||
CertFile: "", |
|
||||
KeyFile: "", |
|
||||
Debug: false, |
|
||||
} |
|
||||
for _, o := range opts { |
|
||||
o(options) |
|
||||
} |
|
||||
if options.Log == nil { |
|
||||
options.Log = log.NewTurnlog(options.Debug, log.Clone("sys.gin", 3)) |
|
||||
} |
|
||||
return |
|
||||
} |
|
||||
@ -1,21 +0,0 @@ |
|||||
package render |
|
||||
|
|
||||
import "net/http" |
|
||||
|
|
||||
// Data contains ContentType and bytes data.
|
|
||||
type Data struct { |
|
||||
ContentType string |
|
||||
Data []byte |
|
||||
} |
|
||||
|
|
||||
// Render (Data) writes data with custom ContentType.
|
|
||||
func (r Data) Render(w http.ResponseWriter) (err error) { |
|
||||
r.WriteContentType(w) |
|
||||
_, err = w.Write(r.Data) |
|
||||
return |
|
||||
} |
|
||||
|
|
||||
// WriteContentType (Data) writes custom ContentType.
|
|
||||
func (r Data) WriteContentType(w http.ResponseWriter) { |
|
||||
writeContentType(w, []string{r.ContentType}) |
|
||||
} |
|
||||
@ -1,84 +0,0 @@ |
|||||
package render |
|
||||
|
|
||||
import ( |
|
||||
"html/template" |
|
||||
"net/http" |
|
||||
) |
|
||||
|
|
||||
var htmlContentType = []string{"text/html; charset=utf-8"} |
|
||||
|
|
||||
/* |
|
||||
Delims 表示一组用于 HTML 模板渲染的左右分隔符。 |
|
||||
*/ |
|
||||
type Delims struct { |
|
||||
// Left delimiter, defaults to {{.
|
|
||||
Left string |
|
||||
// Right delimiter, defaults to }}.
|
|
||||
Right string |
|
||||
} |
|
||||
|
|
||||
/* |
|
||||
包含模板引用及其分隔符 |
|
||||
*/ |
|
||||
type HTMLRender interface { |
|
||||
Instance(string, interface{}) Render |
|
||||
} |
|
||||
type HTMLProduction struct { |
|
||||
Template *template.Template |
|
||||
Delims Delims |
|
||||
} |
|
||||
|
|
||||
func (r HTMLProduction) Instance(name string, data interface{}) Render { |
|
||||
return HTML{ |
|
||||
Template: r.Template, |
|
||||
Name: name, |
|
||||
Data: data, |
|
||||
} |
|
||||
} |
|
||||
|
|
||||
type HTMLDebug struct { |
|
||||
Files []string |
|
||||
Glob string |
|
||||
Delims Delims |
|
||||
FuncMap template.FuncMap |
|
||||
} |
|
||||
|
|
||||
// 返回一个实现Render接口的HTML实例。
|
|
||||
func (r HTMLDebug) Instance(name string, data interface{}) Render { |
|
||||
return HTML{ |
|
||||
Template: r.loadTemplate(), |
|
||||
Name: name, |
|
||||
Data: data, |
|
||||
} |
|
||||
} |
|
||||
func (r HTMLDebug) loadTemplate() *template.Template { |
|
||||
if r.FuncMap == nil { |
|
||||
r.FuncMap = template.FuncMap{} |
|
||||
} |
|
||||
if len(r.Files) > 0 { |
|
||||
return template.Must(template.New("").Delims(r.Delims.Left, r.Delims.Right).Funcs(r.FuncMap).ParseFiles(r.Files...)) |
|
||||
} |
|
||||
if r.Glob != "" { |
|
||||
return template.Must(template.New("").Delims(r.Delims.Left, r.Delims.Right).Funcs(r.FuncMap).ParseGlob(r.Glob)) |
|
||||
} |
|
||||
panic("the HTML debug render was created without files or glob pattern") |
|
||||
} |
|
||||
|
|
||||
type HTML struct { |
|
||||
Template *template.Template |
|
||||
Name string |
|
||||
Data interface{} |
|
||||
} |
|
||||
|
|
||||
func (r HTML) Render(w http.ResponseWriter) error { |
|
||||
r.WriteContentType(w) |
|
||||
|
|
||||
if r.Name == "" { |
|
||||
return r.Template.Execute(w, r.Data) |
|
||||
} |
|
||||
return r.Template.ExecuteTemplate(w, r.Name, r.Data) |
|
||||
} |
|
||||
|
|
||||
func (r HTML) WriteContentType(w http.ResponseWriter) { |
|
||||
writeContentType(w, htmlContentType) |
|
||||
} |
|
||||
@ -1,190 +0,0 @@ |
|||||
package render |
|
||||
|
|
||||
import ( |
|
||||
"bytes" |
|
||||
"fmt" |
|
||||
"html/template" |
|
||||
"net/http" |
|
||||
|
|
||||
json "github.com/json-iterator/go" |
|
||||
|
|
||||
"yunyan/lego/utils" |
|
||||
) |
|
||||
|
|
||||
// JSON contains the given interface object.
|
|
||||
type JSON struct { |
|
||||
Data interface{} |
|
||||
} |
|
||||
|
|
||||
// IndentedJSON contains the given interface object.
|
|
||||
type IndentedJSON struct { |
|
||||
Data interface{} |
|
||||
} |
|
||||
|
|
||||
// SecureJSON contains the given interface object and its prefix.
|
|
||||
type SecureJSON struct { |
|
||||
Prefix string |
|
||||
Data interface{} |
|
||||
} |
|
||||
|
|
||||
// JsonpJSON contains the given interface object its callback.
|
|
||||
type JsonpJSON struct { |
|
||||
Callback string |
|
||||
Data interface{} |
|
||||
} |
|
||||
|
|
||||
// AsciiJSON contains the given interface object.
|
|
||||
type AsciiJSON struct { |
|
||||
Data interface{} |
|
||||
} |
|
||||
|
|
||||
// PureJSON contains the given interface object.
|
|
||||
type PureJSON struct { |
|
||||
Data interface{} |
|
||||
} |
|
||||
|
|
||||
var ( |
|
||||
jsonContentType = []string{"application/json; charset=utf-8"} |
|
||||
jsonpContentType = []string{"application/javascript; charset=utf-8"} |
|
||||
jsonASCIIContentType = []string{"application/json"} |
|
||||
) |
|
||||
|
|
||||
// Render (JSON) writes data with custom ContentType.
|
|
||||
func (r JSON) Render(w http.ResponseWriter) (err error) { |
|
||||
if err = WriteJSON(w, r.Data); err != nil { |
|
||||
panic(err) |
|
||||
} |
|
||||
return |
|
||||
} |
|
||||
|
|
||||
// WriteContentType (JSON) writes JSON ContentType.
|
|
||||
func (r JSON) WriteContentType(w http.ResponseWriter) { |
|
||||
writeContentType(w, jsonContentType) |
|
||||
} |
|
||||
|
|
||||
// WriteJSON marshals the given interface object and writes it with custom ContentType.
|
|
||||
func WriteJSON(w http.ResponseWriter, obj interface{}) error { |
|
||||
writeContentType(w, jsonContentType) |
|
||||
jsonBytes, err := json.Marshal(obj) |
|
||||
if err != nil { |
|
||||
return err |
|
||||
} |
|
||||
_, err = w.Write(jsonBytes) |
|
||||
return err |
|
||||
} |
|
||||
|
|
||||
// Render (IndentedJSON) marshals the given interface object and writes it with custom ContentType.
|
|
||||
func (r IndentedJSON) Render(w http.ResponseWriter) error { |
|
||||
r.WriteContentType(w) |
|
||||
jsonBytes, err := json.MarshalIndent(r.Data, "", " ") |
|
||||
if err != nil { |
|
||||
return err |
|
||||
} |
|
||||
_, err = w.Write(jsonBytes) |
|
||||
return err |
|
||||
} |
|
||||
|
|
||||
// WriteContentType (IndentedJSON) writes JSON ContentType.
|
|
||||
func (r IndentedJSON) WriteContentType(w http.ResponseWriter) { |
|
||||
writeContentType(w, jsonContentType) |
|
||||
} |
|
||||
|
|
||||
// Render (SecureJSON) marshals the given interface object and writes it with custom ContentType.
|
|
||||
func (r SecureJSON) Render(w http.ResponseWriter) error { |
|
||||
r.WriteContentType(w) |
|
||||
jsonBytes, err := json.Marshal(r.Data) |
|
||||
if err != nil { |
|
||||
return err |
|
||||
} |
|
||||
// if the jsonBytes is array values
|
|
||||
if bytes.HasPrefix(jsonBytes, utils.StringToBytes("[")) && bytes.HasSuffix(jsonBytes, |
|
||||
utils.StringToBytes("]")) { |
|
||||
if _, err = w.Write(utils.StringToBytes(r.Prefix)); err != nil { |
|
||||
return err |
|
||||
} |
|
||||
} |
|
||||
_, err = w.Write(jsonBytes) |
|
||||
return err |
|
||||
} |
|
||||
|
|
||||
// WriteContentType (SecureJSON) writes JSON ContentType.
|
|
||||
func (r SecureJSON) WriteContentType(w http.ResponseWriter) { |
|
||||
writeContentType(w, jsonContentType) |
|
||||
} |
|
||||
|
|
||||
// Render (JsonpJSON) marshals the given interface object and writes it and its callback with custom ContentType.
|
|
||||
func (r JsonpJSON) Render(w http.ResponseWriter) (err error) { |
|
||||
r.WriteContentType(w) |
|
||||
ret, err := json.Marshal(r.Data) |
|
||||
if err != nil { |
|
||||
return err |
|
||||
} |
|
||||
|
|
||||
if r.Callback == "" { |
|
||||
_, err = w.Write(ret) |
|
||||
return err |
|
||||
} |
|
||||
|
|
||||
callback := template.JSEscapeString(r.Callback) |
|
||||
if _, err = w.Write(utils.StringToBytes(callback)); err != nil { |
|
||||
return err |
|
||||
} |
|
||||
|
|
||||
if _, err = w.Write(utils.StringToBytes("(")); err != nil { |
|
||||
return err |
|
||||
} |
|
||||
|
|
||||
if _, err = w.Write(ret); err != nil { |
|
||||
return err |
|
||||
} |
|
||||
|
|
||||
if _, err = w.Write(utils.StringToBytes(");")); err != nil { |
|
||||
return err |
|
||||
} |
|
||||
|
|
||||
return nil |
|
||||
} |
|
||||
|
|
||||
// WriteContentType (JsonpJSON) writes Javascript ContentType.
|
|
||||
func (r JsonpJSON) WriteContentType(w http.ResponseWriter) { |
|
||||
writeContentType(w, jsonpContentType) |
|
||||
} |
|
||||
|
|
||||
// Render (AsciiJSON) marshals the given interface object and writes it with custom ContentType.
|
|
||||
func (r AsciiJSON) Render(w http.ResponseWriter) (err error) { |
|
||||
r.WriteContentType(w) |
|
||||
ret, err := json.Marshal(r.Data) |
|
||||
if err != nil { |
|
||||
return err |
|
||||
} |
|
||||
|
|
||||
var buffer bytes.Buffer |
|
||||
for _, r := range utils.BytesToString(ret) { |
|
||||
cvt := string(r) |
|
||||
if r >= 128 { |
|
||||
cvt = fmt.Sprintf("\\u%04x", int64(r)) |
|
||||
} |
|
||||
buffer.WriteString(cvt) |
|
||||
} |
|
||||
|
|
||||
_, err = w.Write(buffer.Bytes()) |
|
||||
return err |
|
||||
} |
|
||||
|
|
||||
// WriteContentType (AsciiJSON) writes JSON ContentType.
|
|
||||
func (r AsciiJSON) WriteContentType(w http.ResponseWriter) { |
|
||||
writeContentType(w, jsonASCIIContentType) |
|
||||
} |
|
||||
|
|
||||
// Render (PureJSON) writes custom ContentType and encodes the given interface object.
|
|
||||
func (r PureJSON) Render(w http.ResponseWriter) error { |
|
||||
r.WriteContentType(w) |
|
||||
encoder := json.NewEncoder(w) |
|
||||
encoder.SetEscapeHTML(false) |
|
||||
return encoder.Encode(r.Data) |
|
||||
} |
|
||||
|
|
||||
// WriteContentType (PureJSON) writes custom ContentType.
|
|
||||
func (r PureJSON) WriteContentType(w http.ResponseWriter) { |
|
||||
writeContentType(w, jsonContentType) |
|
||||
} |
|
||||
@ -1,32 +0,0 @@ |
|||||
package render |
|
||||
|
|
||||
import ( |
|
||||
"net/http" |
|
||||
|
|
||||
"google.golang.org/protobuf/proto" |
|
||||
) |
|
||||
|
|
||||
// ProtoBuf contains the given interface object.
|
|
||||
type ProtoBuf struct { |
|
||||
Data interface{} |
|
||||
} |
|
||||
|
|
||||
var protobufContentType = []string{"application/x-protobuf"} |
|
||||
|
|
||||
// Render (ProtoBuf) marshals the given interface object and writes data with custom ContentType.
|
|
||||
func (r ProtoBuf) Render(w http.ResponseWriter) error { |
|
||||
r.WriteContentType(w) |
|
||||
|
|
||||
bytes, err := proto.Marshal(r.Data.(proto.Message)) |
|
||||
if err != nil { |
|
||||
return err |
|
||||
} |
|
||||
|
|
||||
_, err = w.Write(bytes) |
|
||||
return err |
|
||||
} |
|
||||
|
|
||||
// WriteContentType (ProtoBuf) writes ProtoBuf ContentType.
|
|
||||
func (r ProtoBuf) WriteContentType(w http.ResponseWriter) { |
|
||||
writeContentType(w, protobufContentType) |
|
||||
} |
|
||||
@ -1,44 +0,0 @@ |
|||||
package render |
|
||||
|
|
||||
import ( |
|
||||
"io" |
|
||||
"net/http" |
|
||||
"strconv" |
|
||||
) |
|
||||
|
|
||||
// Reader contains the IO reader and its length, and custom ContentType and other headers.
|
|
||||
type Reader struct { |
|
||||
ContentType string |
|
||||
ContentLength int64 |
|
||||
Reader io.Reader |
|
||||
Headers map[string]string |
|
||||
} |
|
||||
|
|
||||
// Render (Reader) writes data with custom ContentType and headers.
|
|
||||
func (r Reader) Render(w http.ResponseWriter) (err error) { |
|
||||
r.WriteContentType(w) |
|
||||
if r.ContentLength >= 0 { |
|
||||
if r.Headers == nil { |
|
||||
r.Headers = map[string]string{} |
|
||||
} |
|
||||
r.Headers["Content-Length"] = strconv.FormatInt(r.ContentLength, 10) |
|
||||
} |
|
||||
r.writeHeaders(w, r.Headers) |
|
||||
_, err = io.Copy(w, r.Reader) |
|
||||
return |
|
||||
} |
|
||||
|
|
||||
// WriteContentType (Reader) writes custom ContentType.
|
|
||||
func (r Reader) WriteContentType(w http.ResponseWriter) { |
|
||||
writeContentType(w, []string{r.ContentType}) |
|
||||
} |
|
||||
|
|
||||
// writeHeaders writes custom Header.
|
|
||||
func (r Reader) writeHeaders(w http.ResponseWriter, headers map[string]string) { |
|
||||
header := w.Header() |
|
||||
for k, v := range headers { |
|
||||
if header.Get(k) == "" { |
|
||||
header.Set(k, v) |
|
||||
} |
|
||||
} |
|
||||
} |
|
||||
@ -1,25 +0,0 @@ |
|||||
package render |
|
||||
|
|
||||
import ( |
|
||||
"fmt" |
|
||||
"net/http" |
|
||||
) |
|
||||
|
|
||||
// Redirect contains the http request reference and redirects status code and location.
|
|
||||
type Redirect struct { |
|
||||
Code int |
|
||||
Request *http.Request |
|
||||
Location string |
|
||||
} |
|
||||
|
|
||||
// Render (Redirect) redirects the http request to new location and writes redirect response.
|
|
||||
func (r Redirect) Render(w http.ResponseWriter) error { |
|
||||
if (r.Code < http.StatusMultipleChoices || r.Code > http.StatusPermanentRedirect) && r.Code != http.StatusCreated { |
|
||||
panic(fmt.Sprintf("Cannot redirect with status code %d", r.Code)) |
|
||||
} |
|
||||
http.Redirect(w, r.Request, r.Location, r.Code) |
|
||||
return nil |
|
||||
} |
|
||||
|
|
||||
// WriteContentType (Redirect) don't write any ContentType.
|
|
||||
func (r Redirect) WriteContentType(http.ResponseWriter) {} |
|
||||
@ -1,19 +0,0 @@ |
|||||
package render |
|
||||
|
|
||||
import ( |
|
||||
"net/http" |
|
||||
) |
|
||||
|
|
||||
type Render interface { |
|
||||
// Render 使用自定义 ContentType 写入数据。
|
|
||||
Render(http.ResponseWriter) error |
|
||||
// WriteContentType 写入自定义 ContentType。
|
|
||||
WriteContentType(w http.ResponseWriter) |
|
||||
} |
|
||||
|
|
||||
func writeContentType(w http.ResponseWriter, value []string) { |
|
||||
header := w.Header() |
|
||||
if val := header["Content-Type"]; len(val) == 0 { |
|
||||
header["Content-Type"] = value |
|
||||
} |
|
||||
} |
|
||||
@ -1,37 +0,0 @@ |
|||||
package render |
|
||||
|
|
||||
import ( |
|
||||
"fmt" |
|
||||
"net/http" |
|
||||
|
|
||||
"yunyan/lego/utils" |
|
||||
) |
|
||||
|
|
||||
// String contains the given interface object slice and its format.
|
|
||||
type String struct { |
|
||||
Format string |
|
||||
Data []interface{} |
|
||||
} |
|
||||
|
|
||||
var plainContentType = []string{"text/plain; charset=utf-8"} |
|
||||
|
|
||||
// Render (String) writes data with custom ContentType.
|
|
||||
func (r String) Render(w http.ResponseWriter) error { |
|
||||
return WriteString(w, r.Format, r.Data) |
|
||||
} |
|
||||
|
|
||||
// WriteContentType (String) writes Plain ContentType.
|
|
||||
func (r String) WriteContentType(w http.ResponseWriter) { |
|
||||
writeContentType(w, plainContentType) |
|
||||
} |
|
||||
|
|
||||
// WriteString writes data according to its format and write custom ContentType.
|
|
||||
func WriteString(w http.ResponseWriter, format string, data []interface{}) (err error) { |
|
||||
writeContentType(w, plainContentType) |
|
||||
if len(data) > 0 { |
|
||||
_, err = fmt.Fprintf(w, format, data...) |
|
||||
return |
|
||||
} |
|
||||
_, err = w.Write(utils.StringToBytes(format)) |
|
||||
return |
|
||||
} |
|
||||
@ -1,24 +0,0 @@ |
|||||
package render |
|
||||
|
|
||||
import ( |
|
||||
"encoding/xml" |
|
||||
"net/http" |
|
||||
) |
|
||||
|
|
||||
// XML contains the given interface object.
|
|
||||
type XML struct { |
|
||||
Data interface{} |
|
||||
} |
|
||||
|
|
||||
var xmlContentType = []string{"application/xml; charset=utf-8"} |
|
||||
|
|
||||
// Render (XML) encodes the given interface object and writes data with custom ContentType.
|
|
||||
func (r XML) Render(w http.ResponseWriter) error { |
|
||||
r.WriteContentType(w) |
|
||||
return xml.NewEncoder(w).Encode(r.Data) |
|
||||
} |
|
||||
|
|
||||
// WriteContentType (XML) writes XML ContentType for response.
|
|
||||
func (r XML) WriteContentType(w http.ResponseWriter) { |
|
||||
writeContentType(w, xmlContentType) |
|
||||
} |
|
||||
@ -1,32 +0,0 @@ |
|||||
package render |
|
||||
|
|
||||
import ( |
|
||||
"net/http" |
|
||||
|
|
||||
"gopkg.in/yaml.v2" |
|
||||
) |
|
||||
|
|
||||
// YAML contains the given interface object.
|
|
||||
type YAML struct { |
|
||||
Data interface{} |
|
||||
} |
|
||||
|
|
||||
var yamlContentType = []string{"application/x-yaml; charset=utf-8"} |
|
||||
|
|
||||
// Render (YAML) marshals the given interface object and writes data with custom ContentType.
|
|
||||
func (r YAML) Render(w http.ResponseWriter) error { |
|
||||
r.WriteContentType(w) |
|
||||
|
|
||||
bytes, err := yaml.Marshal(r.Data) |
|
||||
if err != nil { |
|
||||
return err |
|
||||
} |
|
||||
|
|
||||
_, err = w.Write(bytes) |
|
||||
return err |
|
||||
} |
|
||||
|
|
||||
// WriteContentType (YAML) writes YAML ContentType for response.
|
|
||||
func (r YAML) WriteContentType(w http.ResponseWriter) { |
|
||||
writeContentType(w, yamlContentType) |
|
||||
} |
|
||||
@ -1,52 +0,0 @@ |
|||||
package gin_test |
|
||||
|
|
||||
import ( |
|
||||
"fmt" |
|
||||
"net/http" |
|
||||
"os" |
|
||||
"os/signal" |
|
||||
"syscall" |
|
||||
"testing" |
|
||||
|
|
||||
"yunyan/lego/sys/gin" |
|
||||
"yunyan/lego/sys/gin/engine" |
|
||||
"yunyan/lego/sys/log" |
|
||||
) |
|
||||
|
|
||||
func Test_sys(t *testing.T) { |
|
||||
if err := log.OnInit(nil, |
|
||||
log.SetFileName("log.log"), |
|
||||
log.SetIsDebug(false), |
|
||||
log.SetEncoder(log.TextEncoder), |
|
||||
); err != nil { |
|
||||
fmt.Printf("log init err:%v", err) |
|
||||
return |
|
||||
} |
|
||||
if sys, err := gin.NewSys(); err != nil { |
|
||||
fmt.Printf("gin init err:%v", err) |
|
||||
} else { |
|
||||
sys.GET("/test", func(c *engine.Context) { |
|
||||
c.JSON(http.StatusOK, "hello") |
|
||||
}) |
|
||||
} |
|
||||
|
|
||||
//监听外部关闭服务信号
|
|
||||
c := make(chan os.Signal, 1) |
|
||||
//添加进程结束信号
|
|
||||
signal.Notify(c, |
|
||||
os.Interrupt, //退出信号 ctrl+c退出
|
|
||||
syscall.SIGHUP, //终端控制进程结束(终端连接断开)
|
|
||||
syscall.SIGINT, //用户发送INTR字符(Ctrl+C)触发
|
|
||||
syscall.SIGTERM, //结束程序(可以被捕获、阻塞或忽略)
|
|
||||
syscall.SIGQUIT) //用户发送QUIT字符(Ctrl+/)触发
|
|
||||
select { |
|
||||
case sig := <-c: |
|
||||
fmt.Println("关闭 signal\n", sig) |
|
||||
} |
|
||||
} |
|
||||
|
|
||||
// /测试签名
|
|
||||
func Test_ParamSign(t *testing.T) { |
|
||||
origin, sgin := gin.ParamSign("@234%67g12q4*67m12#4l67!", map[string]interface{}{"images": []string{"测试资源.png", "11.jpg"}}) |
|
||||
fmt.Println(origin, sgin) |
|
||||
} |
|
||||
@ -1,31 +0,0 @@ |
|||||
package lghttp |
|
||||
|
|
||||
import "context" |
|
||||
|
|
||||
/* |
|
||||
系统描述:mgo数据库驱动系统 |
|
||||
*/ |
|
||||
type ( |
|
||||
ISys interface { |
|
||||
RequestForJson(ctx context.Context, method string, url string, args, result interface{}) (err error) |
|
||||
} |
|
||||
) |
|
||||
|
|
||||
var ( |
|
||||
defsys ISys |
|
||||
) |
|
||||
|
|
||||
func OnInit(config map[string]interface{}, option ...Option) (err error) { |
|
||||
defsys, err = newSys(newOptions(config, option...)) |
|
||||
return |
|
||||
} |
|
||||
|
|
||||
func NewSys(option ...Option) (sys ISys, err error) { |
|
||||
sys, err = newSys(newOptionsByOption(option...)) |
|
||||
return |
|
||||
} |
|
||||
|
|
||||
//请求http
|
|
||||
func RequestForJson(ctx context.Context, method string, url string, args, result interface{}) (err error) { |
|
||||
return defsys.RequestForJson(ctx, method, url, args, result) |
|
||||
} |
|
||||
@ -1,59 +0,0 @@ |
|||||
package lghttp |
|
||||
|
|
||||
import ( |
|
||||
"bytes" |
|
||||
"context" |
|
||||
"yunyan/lego/sys/log" |
|
||||
"encoding/json" |
|
||||
"fmt" |
|
||||
"io" |
|
||||
"net/http" |
|
||||
"time" |
|
||||
) |
|
||||
|
|
||||
func newSys(options Options) (sys *Fasthttp, err error) { |
|
||||
sys = &Fasthttp{options: options} |
|
||||
sys.client = &http.Client{ |
|
||||
Transport: &http.Transport{ |
|
||||
MaxIdleConns: 10, |
|
||||
MaxIdleConnsPerHost: 10, |
|
||||
IdleConnTimeout: 30 * time.Second, |
|
||||
}, |
|
||||
} |
|
||||
return |
|
||||
} |
|
||||
|
|
||||
type Fasthttp struct { |
|
||||
options Options |
|
||||
client *http.Client |
|
||||
} |
|
||||
|
|
||||
func (this *Fasthttp) RequestForJson(ctx context.Context, method string, url string, args, result interface{}) (err error) { |
|
||||
var ( |
|
||||
req *http.Request |
|
||||
resp *http.Response |
|
||||
reqbody []byte |
|
||||
respbody []byte |
|
||||
) |
|
||||
this.options.Log.Debug("RequestForJson", log.Field{Key: "method", Value: method}, log.Field{Key: "url", Value: url}, log.Field{Key: "args", Value: args}) |
|
||||
if reqbody, err = json.Marshal(args); err != nil { |
|
||||
return |
|
||||
} |
|
||||
if req, err = http.NewRequestWithContext(ctx, method, url, bytes.NewBuffer(reqbody)); err != nil { |
|
||||
return |
|
||||
} |
|
||||
req.Header.Set("Content-Type", "application/json") |
|
||||
if resp, err = this.client.Do(req); err != nil { |
|
||||
return |
|
||||
} |
|
||||
defer resp.Body.Close() |
|
||||
if resp.StatusCode != 200 { |
|
||||
err = fmt.Errorf("req fail StatusCode:%d", resp.StatusCode) |
|
||||
} else { |
|
||||
if respbody, err = io.ReadAll(resp.Body); err != nil { |
|
||||
return |
|
||||
} |
|
||||
err = json.Unmarshal(respbody, result) |
|
||||
} |
|
||||
return |
|
||||
} |
|
||||
@ -1,45 +0,0 @@ |
|||||
package lghttp |
|
||||
|
|
||||
import ( |
|
||||
"yunyan/lego/sys/log" |
|
||||
"yunyan/lego/utils/mapstructure" |
|
||||
) |
|
||||
|
|
||||
type Option func(*Options) |
|
||||
type Options struct { |
|
||||
Debug bool //日志是否开启
|
|
||||
Log log.ILogger |
|
||||
MaxConnsPerHost int // 控制每个主机的最大连接数
|
|
||||
ReqTimeOut int // 请求超时时间 单位秒
|
|
||||
} |
|
||||
|
|
||||
func newOptions(config map[string]interface{}, opts ...Option) Options { |
|
||||
options := Options{ |
|
||||
MaxConnsPerHost: 100, |
|
||||
ReqTimeOut: 3, |
|
||||
} |
|
||||
if config != nil { |
|
||||
mapstructure.Decode(config, &options) |
|
||||
} |
|
||||
for _, o := range opts { |
|
||||
o(&options) |
|
||||
} |
|
||||
if options.Log == nil { |
|
||||
options.Log = log.NewTurnlog(options.Debug, log.Clone("sys.fasthttp", 3)) |
|
||||
} |
|
||||
return options |
|
||||
} |
|
||||
|
|
||||
func newOptionsByOption(opts ...Option) Options { |
|
||||
options := Options{ |
|
||||
MaxConnsPerHost: 100, |
|
||||
ReqTimeOut: 3, |
|
||||
} |
|
||||
for _, o := range opts { |
|
||||
o(&options) |
|
||||
} |
|
||||
if options.Log == nil { |
|
||||
options.Log = log.NewTurnlog(options.Debug, log.Clone("sys.fasthttp", 3)) |
|
||||
} |
|
||||
return options |
|
||||
} |
|
||||
@ -1 +0,0 @@ |
|||||
package lghttp_test |
|
||||
@ -1,34 +0,0 @@ |
|||||
package tos |
|
||||
|
|
||||
import ( |
|
||||
"context" |
|
||||
"io" |
|
||||
|
|
||||
"github.com/volcengine/ve-tos-golang-sdk/v2/tos" |
|
||||
) |
|
||||
|
|
||||
type ( |
|
||||
ISys interface { |
|
||||
Get(ctx context.Context, name string, listener tos.DataTransferListener) (resp *tos.GetObjectV2Output, err error) |
|
||||
Put(ctx context.Context, name string, r io.Reader) (err error) |
|
||||
} |
|
||||
) |
|
||||
|
|
||||
var defsys ISys |
|
||||
|
|
||||
func OnInit(config map[string]interface{}, option ...Option) (err error) { |
|
||||
defsys, err = newSys(newOptions(config, option...)) |
|
||||
return |
|
||||
} |
|
||||
|
|
||||
func NewSys(option ...Option) (sys ISys, err error) { |
|
||||
sys, err = newSys(newOptionsByOption(option...)) |
|
||||
return |
|
||||
} |
|
||||
|
|
||||
func Get(ctx context.Context, name string, listener tos.DataTransferListener) (resp *tos.GetObjectV2Output, err error) { |
|
||||
return defsys.Get(ctx, name, listener) |
|
||||
} |
|
||||
func Put(ctx context.Context, name string, r io.Reader) (err error) { |
|
||||
return defsys.Put(ctx, name, r) |
|
||||
} |
|
||||
@ -1,64 +0,0 @@ |
|||||
package tos |
|
||||
|
|
||||
import ( |
|
||||
"yunyan/lego/sys/log" |
|
||||
"yunyan/lego/utils/mapstructure" |
|
||||
) |
|
||||
|
|
||||
type Option func(*Options) |
|
||||
type Options struct { |
|
||||
Debug bool //日志是否开启
|
|
||||
Log log.ILogger |
|
||||
AsccessKey string |
|
||||
SecretKey string |
|
||||
Endpoint string |
|
||||
Region string |
|
||||
BucketName string |
|
||||
} |
|
||||
|
|
||||
func SetAsccessKey(v string) Option { |
|
||||
return func(o *Options) { |
|
||||
o.AsccessKey = v |
|
||||
} |
|
||||
} |
|
||||
func SetSecretKey(v string) Option { |
|
||||
return func(o *Options) { |
|
||||
o.SecretKey = v |
|
||||
} |
|
||||
} |
|
||||
func SetRegion(v string) Option { |
|
||||
return func(o *Options) { |
|
||||
o.Region = v |
|
||||
} |
|
||||
} |
|
||||
|
|
||||
func SetBucketName(v string) Option { |
|
||||
return func(o *Options) { |
|
||||
o.BucketName = v |
|
||||
} |
|
||||
} |
|
||||
|
|
||||
func newOptions(config map[string]interface{}, opts ...Option) Options { |
|
||||
options := Options{} |
|
||||
if config != nil { |
|
||||
mapstructure.Decode(config, &options) |
|
||||
} |
|
||||
for _, o := range opts { |
|
||||
o(&options) |
|
||||
} |
|
||||
if options.Log == nil { |
|
||||
options.Log = log.NewTurnlog(options.Debug, log.Clone("sys.tavily", 3)) |
|
||||
} |
|
||||
return options |
|
||||
} |
|
||||
|
|
||||
func newOptionsByOption(opts ...Option) Options { |
|
||||
options := Options{} |
|
||||
for _, o := range opts { |
|
||||
o(&options) |
|
||||
} |
|
||||
if options.Log == nil { |
|
||||
options.Log = log.NewTurnlog(options.Debug, log.Clone("sys.tavily", 3)) |
|
||||
} |
|
||||
return options |
|
||||
} |
|
||||
@ -1,29 +0,0 @@ |
|||||
package tos_test |
|
||||
|
|
||||
import ( |
|
||||
"context" |
|
||||
"yunyan/lego/sys/sdk/bytedance/tos" |
|
||||
"fmt" |
|
||||
"os" |
|
||||
"testing" |
|
||||
) |
|
||||
|
|
||||
func Test_Sys(t *testing.T) { |
|
||||
if err := tos.OnInit(nil, |
|
||||
tos.SetAsccessKey("AKLTYjQ4NTVjZjkyMGQ1NGRhNWIwODA4YmNmZmU0ZmYwYTg"), |
|
||||
tos.SetSecretKey("TXprNFpEQTFZV1JpTTJGbE5HRTFZemcwWXpFNVptSTBOVGN6TVRnM05XTQ=="), |
|
||||
tos.SetRegion("cn-shanghai"), |
|
||||
tos.SetBucketName("insightcube"), |
|
||||
); err != nil { |
|
||||
return |
|
||||
} else { |
|
||||
var file *os.File |
|
||||
defer file.Close() |
|
||||
if file, err = os.Open("./liwei.text"); err != nil { |
|
||||
fmt.Println("打开文件失败!", err) |
|
||||
return |
|
||||
} |
|
||||
err = tos.Put(context.Background(), "./liwei.text", file) |
|
||||
fmt.Printf("results:%v", err) |
|
||||
} |
|
||||
} |
|
||||
@ -1,46 +0,0 @@ |
|||||
package tos |
|
||||
|
|
||||
import ( |
|
||||
"context" |
|
||||
"io" |
|
||||
|
|
||||
"github.com/volcengine/ve-tos-golang-sdk/v2/tos" |
|
||||
) |
|
||||
|
|
||||
func newSys(options Options) (sys *TOS, err error) { |
|
||||
sys = &TOS{ |
|
||||
options: options, |
|
||||
} |
|
||||
sys.client, err = tos.NewClientV2(options.Endpoint, tos.WithRegion(options.Region), |
|
||||
tos.WithCredentials(tos.NewStaticCredentials(options.AsccessKey, options.SecretKey))) |
|
||||
return |
|
||||
} |
|
||||
|
|
||||
type TOS struct { |
|
||||
options Options |
|
||||
client *tos.ClientV2 |
|
||||
} |
|
||||
|
|
||||
func (this *TOS) Get(ctx context.Context, name string, listener tos.DataTransferListener) (resp *tos.GetObjectV2Output, err error) { |
|
||||
// 下载数据到内存
|
|
||||
resp, err = this.client.GetObjectV2(ctx, &tos.GetObjectV2Input{ |
|
||||
Bucket: this.options.BucketName, |
|
||||
Key: name, |
|
||||
// 获取当前下载进度
|
|
||||
DataTransferListener: listener, |
|
||||
// 下载时重写响应头
|
|
||||
ResponseContentType: "application/json", |
|
||||
}) |
|
||||
return |
|
||||
} |
|
||||
|
|
||||
func (this *TOS) Put(ctx context.Context, name string, r io.Reader) (err error) { |
|
||||
_, err = this.client.PutObjectV2(ctx, &tos.PutObjectV2Input{ |
|
||||
PutObjectBasicInput: tos.PutObjectBasicInput{ |
|
||||
Bucket: this.options.BucketName, |
|
||||
Key: name, |
|
||||
}, |
|
||||
Content: r, |
|
||||
}) |
|
||||
return |
|
||||
} |
|
||||
@ -1,70 +0,0 @@ |
|||||
package timewheel |
|
||||
|
|
||||
import ( |
|
||||
"time" |
|
||||
) |
|
||||
|
|
||||
type ( |
|
||||
ISys interface { |
|
||||
Start() |
|
||||
Stop() |
|
||||
Add(delay time.Duration, handler func(*Task, ...interface{}), args ...interface{}) *Task |
|
||||
AddCron(delay time.Duration, handler func(*Task, ...interface{}), args ...interface{}) *Task |
|
||||
Remove(task *Task) error |
|
||||
NewTimer(delay time.Duration) *Timer |
|
||||
NewTicker(delay time.Duration) *Ticker |
|
||||
AfterFunc(delay time.Duration, callback func()) *Timer |
|
||||
After(delay time.Duration) <-chan time.Time |
|
||||
Sleep(delay time.Duration) |
|
||||
} |
|
||||
) |
|
||||
|
|
||||
var ( |
|
||||
defsys ISys |
|
||||
) |
|
||||
|
|
||||
func OnInit(config map[string]interface{}, option ...Option) (err error) { |
|
||||
if defsys, err = newsys(newOptions(config, option...)); err == nil { |
|
||||
defsys.Start() |
|
||||
} |
|
||||
return |
|
||||
} |
|
||||
|
|
||||
func NewSys(option ...Option) (sys ISys, err error) { |
|
||||
if sys, err = newsys(newOptionsByOption(option...)); err == nil { |
|
||||
sys.Start() |
|
||||
} |
|
||||
return |
|
||||
} |
|
||||
|
|
||||
func Add(delay time.Duration, handler func(*Task, ...interface{}), args ...interface{}) *Task { |
|
||||
return defsys.Add(delay, handler, args...) |
|
||||
} |
|
||||
|
|
||||
func AddCron(delay time.Duration, handler func(*Task, ...interface{}), args ...interface{}) *Task { |
|
||||
return defsys.AddCron(delay, handler, args...) |
|
||||
} |
|
||||
|
|
||||
func Remove(task *Task) error { |
|
||||
return defsys.Remove(task) |
|
||||
} |
|
||||
|
|
||||
func NewTimer(delay time.Duration) *Timer { |
|
||||
return defsys.NewTimer(delay) |
|
||||
} |
|
||||
|
|
||||
func NewTicker(delay time.Duration) *Ticker { |
|
||||
return defsys.NewTicker(delay) |
|
||||
} |
|
||||
|
|
||||
func AfterFunc(delay time.Duration, callback func()) *Timer { |
|
||||
return defsys.AfterFunc(delay, callback) |
|
||||
} |
|
||||
|
|
||||
func After(delay time.Duration) <-chan time.Time { |
|
||||
return defsys.After(delay) |
|
||||
} |
|
||||
|
|
||||
func Sleep(delay time.Duration) { |
|
||||
defsys.Sleep(delay) |
|
||||
} |
|
||||
@ -1,66 +0,0 @@ |
|||||
package timewheel |
|
||||
|
|
||||
import ( |
|
||||
"yunyan/lego/sys/log" |
|
||||
"yunyan/lego/utils/mapstructure" |
|
||||
"time" |
|
||||
) |
|
||||
|
|
||||
type Option func(*Options) |
|
||||
type Options struct { |
|
||||
Tick time.Duration //不小于 10毫秒
|
|
||||
BucketsNum int |
|
||||
} |
|
||||
|
|
||||
func SetTick(v time.Duration) Option { |
|
||||
return func(o *Options) { |
|
||||
o.Tick = v |
|
||||
} |
|
||||
} |
|
||||
|
|
||||
func SetBucketsNum(v int) Option { |
|
||||
return func(o *Options) { |
|
||||
o.BucketsNum = v |
|
||||
} |
|
||||
} |
|
||||
|
|
||||
func newOptions(config map[string]interface{}, opts ...Option) Options { |
|
||||
options := Options{ |
|
||||
Tick: time.Second, |
|
||||
BucketsNum: 1024, |
|
||||
} |
|
||||
if config != nil { |
|
||||
mapstructure.Decode(config, &options) |
|
||||
} |
|
||||
for _, o := range opts { |
|
||||
o(&options) |
|
||||
} |
|
||||
if options.Tick < 100*time.Millisecond { |
|
||||
log.Errorf("创建时间轮参数异常 Tick 必须大于 100 ms ") |
|
||||
options.Tick = 100 * time.Millisecond |
|
||||
} |
|
||||
if options.BucketsNum < 0 { |
|
||||
log.Errorf("创建时间轮参数异常 BucketsNum 必须大于 0 ") |
|
||||
options.BucketsNum = 1 |
|
||||
} |
|
||||
return options |
|
||||
} |
|
||||
|
|
||||
func newOptionsByOption(opts ...Option) Options { |
|
||||
options := Options{ |
|
||||
Tick: 1000, |
|
||||
BucketsNum: 1024, |
|
||||
} |
|
||||
for _, o := range opts { |
|
||||
o(&options) |
|
||||
} |
|
||||
if options.Tick < 100*time.Millisecond { |
|
||||
log.Warnf("创建时间轮参数异常 Tick 必须大于 100 ms ") |
|
||||
options.Tick = 100 * time.Millisecond |
|
||||
} |
|
||||
if options.BucketsNum < 0 { |
|
||||
log.Warnf("创建时间轮参数异常 BucketsNum 必须大于 0 ") |
|
||||
options.BucketsNum = 1 |
|
||||
} |
|
||||
return options |
|
||||
} |
|
||||
@ -1,34 +0,0 @@ |
|||||
package timewheel |
|
||||
|
|
||||
import ( |
|
||||
"sync" |
|
||||
) |
|
||||
|
|
||||
var incr = 0 |
|
||||
|
|
||||
var ( |
|
||||
defaultTaskPool = newTaskPool() |
|
||||
) |
|
||||
|
|
||||
type taskPool struct { |
|
||||
bp *sync.Pool |
|
||||
} |
|
||||
|
|
||||
func newTaskPool() *taskPool { |
|
||||
return &taskPool{ |
|
||||
bp: &sync.Pool{ |
|
||||
New: func() interface{} { |
|
||||
return &Task{} |
|
||||
}, |
|
||||
}, |
|
||||
} |
|
||||
} |
|
||||
|
|
||||
func (pool *taskPool) get() *Task { |
|
||||
return pool.bp.Get().(*Task) |
|
||||
} |
|
||||
|
|
||||
func (pool *taskPool) put(obj *Task) { |
|
||||
obj.Reset() |
|
||||
pool.bp.Put(obj) |
|
||||
} |
|
||||
@ -1,457 +0,0 @@ |
|||||
package timewheel |
|
||||
|
|
||||
import ( |
|
||||
"context" |
|
||||
"yunyan/lego" |
|
||||
"yunyan/lego/sys/log" |
|
||||
"fmt" |
|
||||
"runtime" |
|
||||
"sync" |
|
||||
"sync/atomic" |
|
||||
"time" |
|
||||
) |
|
||||
|
|
||||
// 创建一个时间轮
|
|
||||
func newsys(options Options) (sys *TimeWheel, err error) { |
|
||||
sys = &TimeWheel{ |
|
||||
// tick
|
|
||||
tick: options.Tick, |
|
||||
tickQueue: make(chan time.Time, 10), |
|
||||
|
|
||||
// store
|
|
||||
bucketsNum: options.BucketsNum, |
|
||||
bucketIndexes: make(map[taskID]int, 1024*100), |
|
||||
buckets: make([]map[taskID]*Task, options.BucketsNum), |
|
||||
currentIndex: 0, |
|
||||
|
|
||||
// signal
|
|
||||
addC: make(chan *Task, 1024*5), |
|
||||
removeC: make(chan *Task, 1024*2), |
|
||||
stopC: make(chan struct{}), |
|
||||
} |
|
||||
|
|
||||
for i := 0; i < options.BucketsNum; i++ { |
|
||||
sys.buckets[i] = make(map[taskID]*Task, 16) |
|
||||
} |
|
||||
|
|
||||
return |
|
||||
} |
|
||||
|
|
||||
const ( |
|
||||
typeTimer taskType = iota |
|
||||
typeTicker |
|
||||
|
|
||||
modeIsCircle = true |
|
||||
modeNotCircle = false |
|
||||
|
|
||||
modeIsAsync = true |
|
||||
modeNotAsync = false |
|
||||
) |
|
||||
|
|
||||
type ( |
|
||||
taskType int64 |
|
||||
taskID int64 |
|
||||
Task struct { |
|
||||
delay time.Duration |
|
||||
id taskID |
|
||||
round int |
|
||||
args []interface{} |
|
||||
callback func(*Task, ...interface{}) |
|
||||
async bool //异步执行
|
|
||||
stop bool //是否停止
|
|
||||
circle bool //是否循环
|
|
||||
} |
|
||||
TimeWheel struct { |
|
||||
randomID int64 |
|
||||
tick time.Duration |
|
||||
ticker *time.Ticker |
|
||||
tickQueue chan time.Time |
|
||||
bucketsNum int |
|
||||
buckets []map[taskID]*Task // key: added item, value: *Task
|
|
||||
bucketIndexes map[taskID]int // key: added item, value: bucket position
|
|
||||
currentIndex int |
|
||||
onceStart sync.Once |
|
||||
addC chan *Task |
|
||||
removeC chan *Task |
|
||||
stopC chan struct{} |
|
||||
exited bool |
|
||||
} |
|
||||
) |
|
||||
|
|
||||
// for sync.Pool
|
|
||||
func (t *Task) Reset() { |
|
||||
t.round = 0 |
|
||||
t.callback = nil |
|
||||
|
|
||||
t.async = false |
|
||||
t.stop = false |
|
||||
t.circle = false |
|
||||
} |
|
||||
|
|
||||
// 启动时间轮
|
|
||||
func (this *TimeWheel) Start() { |
|
||||
// onlye once start
|
|
||||
this.onceStart.Do( |
|
||||
func() { |
|
||||
this.ticker = time.NewTicker(this.tick) |
|
||||
go this.schduler() |
|
||||
go this.tickGenerator() |
|
||||
}, |
|
||||
) |
|
||||
} |
|
||||
|
|
||||
func (this *TimeWheel) Add(delay time.Duration, handler func(*Task, ...interface{}), args ...interface{}) *Task { |
|
||||
return this.addAny(delay, modeNotCircle, modeIsAsync, handler, args...) |
|
||||
} |
|
||||
|
|
||||
// AddCron add interval task
|
|
||||
func (this *TimeWheel) AddCron(delay time.Duration, handler func(*Task, ...interface{}), args ...interface{}) *Task { |
|
||||
return this.addAny(delay, true, modeIsAsync, handler, args...) |
|
||||
} |
|
||||
|
|
||||
func (this *TimeWheel) Remove(task *Task) error { |
|
||||
this.removeC <- task |
|
||||
return nil |
|
||||
} |
|
||||
|
|
||||
// 停止时间轮
|
|
||||
func (this *TimeWheel) Stop() { |
|
||||
this.stopC <- struct{}{} |
|
||||
} |
|
||||
|
|
||||
// 此处写法 为监控时间轮是否正常执行
|
|
||||
func (this *TimeWheel) tickGenerator() { |
|
||||
if this.tickQueue == nil { |
|
||||
return |
|
||||
} |
|
||||
|
|
||||
for !this.exited { |
|
||||
select { |
|
||||
case <-this.ticker.C: |
|
||||
select { |
|
||||
case this.tickQueue <- time.Now(): |
|
||||
default: |
|
||||
panic("raise long time blocking") |
|
||||
} |
|
||||
} |
|
||||
} |
|
||||
} |
|
||||
|
|
||||
// 调度器
|
|
||||
func (this *TimeWheel) schduler() { |
|
||||
queue := this.ticker.C |
|
||||
if this.tickQueue != nil { |
|
||||
queue = this.tickQueue |
|
||||
} |
|
||||
|
|
||||
for { |
|
||||
select { |
|
||||
case <-queue: |
|
||||
this.handleTick() |
|
||||
case task := <-this.addC: |
|
||||
this.put(task) |
|
||||
case key := <-this.removeC: |
|
||||
this.remove(key) |
|
||||
case <-this.stopC: |
|
||||
this.exited = true |
|
||||
this.ticker.Stop() |
|
||||
return |
|
||||
} |
|
||||
} |
|
||||
} |
|
||||
|
|
||||
// 清理
|
|
||||
func (this *TimeWheel) collectTask(task *Task) { |
|
||||
if index, ok := this.bucketIndexes[task.id]; ok { |
|
||||
delete(this.bucketIndexes, task.id) |
|
||||
delete(this.buckets[index], task.id) |
|
||||
} |
|
||||
} |
|
||||
|
|
||||
func (this *TimeWheel) recoverTask(task *Task) { |
|
||||
defaultTaskPool.put(task) |
|
||||
} |
|
||||
|
|
||||
func (this *TimeWheel) handleTick() { |
|
||||
bucket := this.buckets[this.currentIndex] |
|
||||
for k, task := range bucket { |
|
||||
if task.stop { |
|
||||
this.collectTask(task) |
|
||||
this.recoverTask(task) |
|
||||
continue |
|
||||
} |
|
||||
|
|
||||
if bucket[k].round > 0 { |
|
||||
bucket[k].round-- |
|
||||
continue |
|
||||
} |
|
||||
this.collectTask(task) |
|
||||
if task.async { |
|
||||
go func(_task *Task) { |
|
||||
defer func() { //程序异常 收集异常信息传递给前端显示
|
|
||||
if r := recover(); r != nil { |
|
||||
buf := make([]byte, 4096) |
|
||||
l := runtime.Stack(buf, false) |
|
||||
err := fmt.Errorf("%v: %s", r, buf[:l]) |
|
||||
log.Errorf("[timewheel] calltask err:%s", err.Error()) |
|
||||
} |
|
||||
}() |
|
||||
this.calltask(_task, _task.args...) |
|
||||
if _task.circle { //重新进入队列中
|
|
||||
// this.putCircle(_task, true)
|
|
||||
this.addC <- _task |
|
||||
} else { |
|
||||
this.recoverTask(_task) |
|
||||
} |
|
||||
}(task) |
|
||||
} else { |
|
||||
this.calltask(task, task.args...) |
|
||||
//循环执行
|
|
||||
if task.circle { |
|
||||
this.putCircle(task, true) |
|
||||
} else { |
|
||||
this.recoverTask(task) |
|
||||
} |
|
||||
} |
|
||||
} |
|
||||
|
|
||||
if this.currentIndex == this.bucketsNum-1 { |
|
||||
this.currentIndex = 0 |
|
||||
return |
|
||||
} |
|
||||
|
|
||||
this.currentIndex++ |
|
||||
} |
|
||||
|
|
||||
// 执行时间轮事件 捕捉异常错误 防止程序崩溃
|
|
||||
func (this *TimeWheel) calltask(task *Task, args ...interface{}) { |
|
||||
defer lego.Recover("TimeWheel") |
|
||||
if task.callback == nil { |
|
||||
log.Error("sys.timeWheel task callback err!", log.Field{Key: "task", Value: task}) |
|
||||
return |
|
||||
} |
|
||||
task.callback(task, task.args...) |
|
||||
} |
|
||||
|
|
||||
func (this *TimeWheel) addAny(delay time.Duration, circle, async bool, callback func(*Task, ...interface{}), agr ...interface{}) *Task { |
|
||||
if delay <= 0 { |
|
||||
delay = this.tick |
|
||||
} |
|
||||
|
|
||||
id := this.genUniqueID() |
|
||||
|
|
||||
var task *Task |
|
||||
task = defaultTaskPool.get() |
|
||||
task.delay = delay |
|
||||
task.id = id |
|
||||
task.args = agr |
|
||||
task.callback = callback |
|
||||
task.circle = circle |
|
||||
task.async = async // refer to src/runtime/time.go
|
|
||||
|
|
||||
this.addC <- task |
|
||||
return task |
|
||||
} |
|
||||
|
|
||||
func (this *TimeWheel) put(task *Task) { |
|
||||
this.store(task, false) |
|
||||
} |
|
||||
|
|
||||
func (this *TimeWheel) putCircle(task *Task, circleMode bool) { |
|
||||
this.store(task, circleMode) |
|
||||
} |
|
||||
|
|
||||
func (this *TimeWheel) store(task *Task, circleMode bool) { |
|
||||
round := this.calculateRound(task.delay) |
|
||||
index := this.calculateIndex(task.delay) |
|
||||
|
|
||||
if round > 0 && circleMode { |
|
||||
task.round = round - 1 |
|
||||
} else { |
|
||||
task.round = round |
|
||||
} |
|
||||
|
|
||||
this.bucketIndexes[task.id] = index |
|
||||
this.buckets[index][task.id] = task |
|
||||
} |
|
||||
|
|
||||
func (this *TimeWheel) calculateRound(delay time.Duration) (round int) { |
|
||||
delaySeconds := delay.Seconds() |
|
||||
tickSeconds := this.tick.Seconds() |
|
||||
round = int(delaySeconds / tickSeconds / float64(this.bucketsNum)) |
|
||||
return |
|
||||
} |
|
||||
|
|
||||
func (this *TimeWheel) calculateIndex(delay time.Duration) (index int) { |
|
||||
delaySeconds := delay.Seconds() |
|
||||
tickSeconds := this.tick.Seconds() |
|
||||
index = (int(float64(this.currentIndex) + delaySeconds/tickSeconds)) % this.bucketsNum |
|
||||
return |
|
||||
} |
|
||||
|
|
||||
func (this *TimeWheel) remove(task *Task) { |
|
||||
this.collectTask(task) |
|
||||
this.recoverTask(task) |
|
||||
} |
|
||||
|
|
||||
func (this *TimeWheel) NewTimer(delay time.Duration) *Timer { |
|
||||
queue := make(chan bool, 1) // buf = 1, refer to src/time/sleep.go
|
|
||||
task := this.addAny(delay, |
|
||||
modeNotCircle, |
|
||||
modeNotAsync, |
|
||||
func(*Task, ...interface{}) { |
|
||||
notfiyChannel(queue) |
|
||||
}, |
|
||||
) |
|
||||
|
|
||||
// init timer
|
|
||||
ctx, cancel := context.WithCancel(context.Background()) |
|
||||
timer := &Timer{ |
|
||||
this: this, |
|
||||
C: queue, // faster
|
|
||||
task: task, |
|
||||
Ctx: ctx, |
|
||||
cancel: cancel, |
|
||||
} |
|
||||
|
|
||||
return timer |
|
||||
} |
|
||||
|
|
||||
func (this *TimeWheel) AfterFunc(delay time.Duration, callback func()) *Timer { |
|
||||
queue := make(chan bool, 1) |
|
||||
task := this.addAny(delay, |
|
||||
modeNotCircle, modeIsAsync, |
|
||||
func(*Task, ...interface{}) { |
|
||||
callback() |
|
||||
notfiyChannel(queue) |
|
||||
}, |
|
||||
) |
|
||||
|
|
||||
// init timer
|
|
||||
ctx, cancel := context.WithCancel(context.Background()) |
|
||||
timer := &Timer{ |
|
||||
this: this, |
|
||||
C: queue, // faster
|
|
||||
task: task, |
|
||||
Ctx: ctx, |
|
||||
cancel: cancel, |
|
||||
fn: callback, |
|
||||
} |
|
||||
|
|
||||
return timer |
|
||||
} |
|
||||
|
|
||||
func (this *TimeWheel) NewTicker(delay time.Duration) *Ticker { |
|
||||
queue := make(chan bool, 1) |
|
||||
task := this.addAny(delay, |
|
||||
modeIsCircle, |
|
||||
modeNotAsync, |
|
||||
func(*Task, ...interface{}) { |
|
||||
notfiyChannel(queue) |
|
||||
}, |
|
||||
) |
|
||||
|
|
||||
// init ticker
|
|
||||
ctx, cancel := context.WithCancel(context.Background()) |
|
||||
ticker := &Ticker{ |
|
||||
task: task, |
|
||||
this: this, |
|
||||
C: queue, |
|
||||
Ctx: ctx, |
|
||||
cancel: cancel, |
|
||||
} |
|
||||
|
|
||||
return ticker |
|
||||
} |
|
||||
|
|
||||
func (this *TimeWheel) After(delay time.Duration) <-chan time.Time { |
|
||||
queue := make(chan time.Time, 1) |
|
||||
this.addAny(delay, |
|
||||
modeNotCircle, modeNotAsync, |
|
||||
func(*Task, ...interface{}) { |
|
||||
queue <- time.Now() |
|
||||
}, |
|
||||
) |
|
||||
return queue |
|
||||
} |
|
||||
|
|
||||
func (this *TimeWheel) Sleep(delay time.Duration) { |
|
||||
queue := make(chan bool, 1) |
|
||||
this.addAny(delay, |
|
||||
modeNotCircle, modeNotAsync, |
|
||||
func(*Task, ...interface{}) { |
|
||||
queue <- true |
|
||||
}, |
|
||||
) |
|
||||
<-queue |
|
||||
} |
|
||||
|
|
||||
// similar to golang std timer
|
|
||||
type Timer struct { |
|
||||
task *Task |
|
||||
this *TimeWheel |
|
||||
fn func() // external custom func
|
|
||||
C chan bool |
|
||||
|
|
||||
cancel context.CancelFunc |
|
||||
Ctx context.Context |
|
||||
} |
|
||||
|
|
||||
func (t *Timer) Reset(delay time.Duration) { |
|
||||
var task *Task |
|
||||
if t.fn != nil { // use AfterFunc
|
|
||||
task = t.this.addAny(delay, |
|
||||
modeNotCircle, modeIsAsync, // must async mode
|
|
||||
func(*Task, ...interface{}) { |
|
||||
t.fn() |
|
||||
notfiyChannel(t.C) |
|
||||
}, |
|
||||
) |
|
||||
} else { |
|
||||
task = t.this.addAny(delay, |
|
||||
modeNotCircle, modeNotAsync, |
|
||||
func(*Task, ...interface{}) { |
|
||||
notfiyChannel(t.C) |
|
||||
}, |
|
||||
) |
|
||||
} |
|
||||
|
|
||||
t.task = task |
|
||||
} |
|
||||
|
|
||||
func (t *Timer) Stop() { |
|
||||
t.task.stop = true |
|
||||
t.cancel() |
|
||||
t.this.Remove(t.task) |
|
||||
} |
|
||||
|
|
||||
func (t *Timer) StopFunc(callback func()) { |
|
||||
t.fn = callback |
|
||||
} |
|
||||
|
|
||||
type Ticker struct { |
|
||||
this *TimeWheel |
|
||||
task *Task |
|
||||
cancel context.CancelFunc |
|
||||
|
|
||||
C chan bool |
|
||||
Ctx context.Context |
|
||||
} |
|
||||
|
|
||||
func (t *Ticker) Stop() { |
|
||||
t.task.stop = true |
|
||||
t.cancel() |
|
||||
t.this.Remove(t.task) |
|
||||
} |
|
||||
|
|
||||
func notfiyChannel(q chan bool) { |
|
||||
select { |
|
||||
case q <- true: |
|
||||
default: |
|
||||
} |
|
||||
} |
|
||||
|
|
||||
func (this *TimeWheel) genUniqueID() taskID { |
|
||||
id := atomic.AddInt64(&this.randomID, 1) |
|
||||
return taskID(id) |
|
||||
} |
|
||||
@ -1,57 +0,0 @@ |
|||||
package timewheel_test |
|
||||
|
|
||||
import ( |
|
||||
"fmt" |
|
||||
"testing" |
|
||||
"time" |
|
||||
|
|
||||
"yunyan/lego/sys/timewheel" |
|
||||
) |
|
||||
|
|
||||
func checkTimeCost(t *testing.T, start, end time.Time, before int, after int) bool { |
|
||||
due := end.Sub(start) |
|
||||
if due > time.Duration(after)*time.Millisecond { |
|
||||
t.Error("delay run") |
|
||||
return false |
|
||||
} |
|
||||
|
|
||||
if due < time.Duration(before)*time.Millisecond { |
|
||||
t.Error("run ahead") |
|
||||
return false |
|
||||
} |
|
||||
|
|
||||
return true |
|
||||
} |
|
||||
|
|
||||
func TestAddFunc(t *testing.T) { |
|
||||
tw, _ := timewheel.NewSys(timewheel.SetTick(100), timewheel.SetBucketsNum(10)) |
|
||||
tw.Start() |
|
||||
defer tw.Stop() |
|
||||
|
|
||||
for index := 1; index < 6; index++ { |
|
||||
queue := make(chan bool, 0) |
|
||||
start := time.Now() |
|
||||
tw.Add(time.Duration(index)*time.Second, func(*timewheel.Task, ...interface{}) { |
|
||||
queue <- true |
|
||||
}) |
|
||||
|
|
||||
<-queue |
|
||||
|
|
||||
before := index*1000 - 200 |
|
||||
after := index*1000 + 200 |
|
||||
checkTimeCost(t, start, time.Now(), before, after) |
|
||||
fmt.Println("time since: ", time.Since(start).String()) |
|
||||
} |
|
||||
} |
|
||||
|
|
||||
func Test_AddFunc(t *testing.T) { |
|
||||
tw, _ := timewheel.NewSys(timewheel.SetTick(100*time.Millisecond), timewheel.SetBucketsNum(10)) |
|
||||
tw.Start() |
|
||||
defer tw.Stop() |
|
||||
start := time.Now() |
|
||||
tw.AddCron(time.Second, func(*timewheel.Task, ...interface{}) { |
|
||||
fmt.Println("time since: ", time.Since(start).String()) |
|
||||
}) |
|
||||
|
|
||||
time.Sleep(10 * time.Second) |
|
||||
} |
|
||||
@ -1,84 +0,0 @@ |
|||||
package container |
|
||||
|
|
||||
import "sync" |
|
||||
|
|
||||
// BeeMap is a map with lock
|
|
||||
|
|
||||
type BeeMap struct { |
|
||||
lock *sync.RWMutex |
|
||||
bm map[interface{}]interface{} |
|
||||
} |
|
||||
|
|
||||
// NewBeeMap return new safemap
|
|
||||
func NewBeeMap() *BeeMap { |
|
||||
return &BeeMap{ |
|
||||
lock: new(sync.RWMutex), |
|
||||
bm: make(map[interface{}]interface{}), |
|
||||
} |
|
||||
} |
|
||||
|
|
||||
// Get from maps return the k's value
|
|
||||
func (m *BeeMap) Get(k interface{}) interface{} { |
|
||||
m.lock.RLock() |
|
||||
if val, ok := m.bm[k]; ok { |
|
||||
m.lock.RUnlock() |
|
||||
return val |
|
||||
} |
|
||||
m.lock.RUnlock() |
|
||||
return nil |
|
||||
} |
|
||||
|
|
||||
// Set Maps the given key and value. Returns false
|
|
||||
// if the key is already in the map and changes nothing.
|
|
||||
func (m *BeeMap) Set(k interface{}, v interface{}) bool { |
|
||||
m.lock.Lock() |
|
||||
if val, ok := m.bm[k]; !ok { |
|
||||
m.bm[k] = v |
|
||||
m.lock.Unlock() |
|
||||
} else if val != v { |
|
||||
m.bm[k] = v |
|
||||
m.lock.Unlock() |
|
||||
} else { |
|
||||
m.lock.Unlock() |
|
||||
return false |
|
||||
} |
|
||||
return true |
|
||||
} |
|
||||
|
|
||||
// Check Returns true if k is exist in the map.
|
|
||||
func (m *BeeMap) Check(k interface{}) bool { |
|
||||
m.lock.RLock() |
|
||||
if _, ok := m.bm[k]; !ok { |
|
||||
m.lock.RUnlock() |
|
||||
return false |
|
||||
} |
|
||||
m.lock.RUnlock() |
|
||||
return true |
|
||||
} |
|
||||
|
|
||||
// Delete the given key and value.
|
|
||||
func (m *BeeMap) Delete(k interface{}) { |
|
||||
m.lock.Lock() |
|
||||
delete(m.bm, k) |
|
||||
m.lock.Unlock() |
|
||||
} |
|
||||
|
|
||||
func (m *BeeMap) DeleteAll() { |
|
||||
m.lock.Lock() |
|
||||
for k, _ := range m.bm { |
|
||||
delete(m.bm, k) |
|
||||
} |
|
||||
|
|
||||
m.lock.Unlock() |
|
||||
} |
|
||||
|
|
||||
// Items returns all items in safemap.
|
|
||||
func (m *BeeMap) Items() map[interface{}]interface{} { |
|
||||
m.lock.RLock() |
|
||||
r := make(map[interface{}]interface{}) |
|
||||
for k, v := range m.bm { |
|
||||
r[k] = v |
|
||||
} |
|
||||
m.lock.RUnlock() |
|
||||
return r |
|
||||
} |
|
||||
@ -1,250 +0,0 @@ |
|||||
package container |
|
||||
|
|
||||
import ( |
|
||||
"sync" |
|
||||
|
|
||||
json "github.com/json-iterator/go" |
|
||||
) |
|
||||
|
|
||||
var SHARD_COUNT = 32 |
|
||||
|
|
||||
type ConcurrentMap []*ConcurrentMapShared |
|
||||
|
|
||||
type ConcurrentMapShared struct { |
|
||||
items map[string]interface{} |
|
||||
sync.RWMutex |
|
||||
} |
|
||||
|
|
||||
func NewConcurrentMap() ConcurrentMap { |
|
||||
m := make(ConcurrentMap, SHARD_COUNT) |
|
||||
for i := 0; i < SHARD_COUNT; i++ { |
|
||||
m[i] = &ConcurrentMapShared{items: make(map[string]interface{})} |
|
||||
} |
|
||||
return m |
|
||||
} |
|
||||
|
|
||||
func (m ConcurrentMap) GetShard(key string) *ConcurrentMapShared { |
|
||||
return m[uint(fnv32(key))%uint(SHARD_COUNT)] |
|
||||
} |
|
||||
func (m ConcurrentMap) MSet(data map[string]interface{}) { |
|
||||
for key, value := range data { |
|
||||
shard := m.GetShard(key) |
|
||||
shard.Lock() |
|
||||
shard.items[key] = value |
|
||||
shard.Unlock() |
|
||||
} |
|
||||
} |
|
||||
func (m ConcurrentMap) Set(key string, value interface{}) { |
|
||||
// Get map shard.
|
|
||||
shard := m.GetShard(key) |
|
||||
shard.Lock() |
|
||||
shard.items[key] = value |
|
||||
shard.Unlock() |
|
||||
} |
|
||||
func (m ConcurrentMap) Upsert(key string, value interface{}, cb UpsertCb) (res interface{}) { |
|
||||
shard := m.GetShard(key) |
|
||||
shard.Lock() |
|
||||
v, ok := shard.items[key] |
|
||||
res = cb(ok, v, value) |
|
||||
shard.items[key] = res |
|
||||
shard.Unlock() |
|
||||
return res |
|
||||
} |
|
||||
func (m ConcurrentMap) SetIfAbsent(key string, value interface{}) bool { |
|
||||
// Get map shard.
|
|
||||
shard := m.GetShard(key) |
|
||||
shard.Lock() |
|
||||
_, ok := shard.items[key] |
|
||||
if !ok { |
|
||||
shard.items[key] = value |
|
||||
} |
|
||||
shard.Unlock() |
|
||||
return !ok |
|
||||
} |
|
||||
func (m ConcurrentMap) Get(key string) (interface{}, bool) { |
|
||||
shard := m.GetShard(key) |
|
||||
shard.RLock() |
|
||||
val, ok := shard.items[key] |
|
||||
shard.RUnlock() |
|
||||
return val, ok |
|
||||
} |
|
||||
func (m ConcurrentMap) Count() int { |
|
||||
count := 0 |
|
||||
for i := 0; i < SHARD_COUNT; i++ { |
|
||||
shard := m[i] |
|
||||
shard.RLock() |
|
||||
count += len(shard.items) |
|
||||
shard.RUnlock() |
|
||||
} |
|
||||
return count |
|
||||
} |
|
||||
func (m ConcurrentMap) Has(key string) bool { |
|
||||
// Get shard
|
|
||||
shard := m.GetShard(key) |
|
||||
shard.RLock() |
|
||||
// See if element is within shard.
|
|
||||
_, ok := shard.items[key] |
|
||||
shard.RUnlock() |
|
||||
return ok |
|
||||
} |
|
||||
func (m ConcurrentMap) Remove(key string) { |
|
||||
// Try to get shard.
|
|
||||
shard := m.GetShard(key) |
|
||||
shard.Lock() |
|
||||
delete(shard.items, key) |
|
||||
shard.Unlock() |
|
||||
} |
|
||||
func (m ConcurrentMap) RemoveCb(key string, cb RemoveCb) bool { |
|
||||
// Try to get shard.
|
|
||||
shard := m.GetShard(key) |
|
||||
shard.Lock() |
|
||||
v, ok := shard.items[key] |
|
||||
remove := cb(key, v, ok) |
|
||||
if remove && ok { |
|
||||
delete(shard.items, key) |
|
||||
} |
|
||||
shard.Unlock() |
|
||||
return remove |
|
||||
} |
|
||||
func (m ConcurrentMap) Pop(key string) (v interface{}, exists bool) { |
|
||||
// Try to get shard.
|
|
||||
shard := m.GetShard(key) |
|
||||
shard.Lock() |
|
||||
v, exists = shard.items[key] |
|
||||
delete(shard.items, key) |
|
||||
shard.Unlock() |
|
||||
return v, exists |
|
||||
} |
|
||||
func (m ConcurrentMap) IsEmpty() bool { |
|
||||
return m.Count() == 0 |
|
||||
} |
|
||||
func (m ConcurrentMap) Iter() <-chan Tuple { |
|
||||
chans := snapshot(m) |
|
||||
ch := make(chan Tuple) |
|
||||
go fanIn(chans, ch) |
|
||||
return ch |
|
||||
} |
|
||||
func (m ConcurrentMap) IterBuffered() <-chan Tuple { |
|
||||
chans := snapshot(m) |
|
||||
total := 0 |
|
||||
for _, c := range chans { |
|
||||
total += cap(c) |
|
||||
} |
|
||||
ch := make(chan Tuple, total) |
|
||||
go fanIn(chans, ch) |
|
||||
return ch |
|
||||
} |
|
||||
func (m ConcurrentMap) Items() map[string]interface{} { |
|
||||
tmp := make(map[string]interface{}) |
|
||||
|
|
||||
// Insert items to temporary map.
|
|
||||
for item := range m.IterBuffered() { |
|
||||
tmp[item.Key] = item.Val |
|
||||
} |
|
||||
|
|
||||
return tmp |
|
||||
} |
|
||||
func (m ConcurrentMap) IterCb(fn IterCb) { |
|
||||
for idx := range m { |
|
||||
shard := (m)[idx] |
|
||||
shard.RLock() |
|
||||
for key, value := range shard.items { |
|
||||
fn(key, value) |
|
||||
} |
|
||||
shard.RUnlock() |
|
||||
} |
|
||||
} |
|
||||
func (m ConcurrentMap) Keys() []string { |
|
||||
count := m.Count() |
|
||||
ch := make(chan string, count) |
|
||||
go func() { |
|
||||
// Foreach shard.
|
|
||||
wg := sync.WaitGroup{} |
|
||||
wg.Add(SHARD_COUNT) |
|
||||
for _, shard := range m { |
|
||||
go func(shard *ConcurrentMapShared) { |
|
||||
// Foreach key, value pair.
|
|
||||
shard.RLock() |
|
||||
for key := range shard.items { |
|
||||
ch <- key |
|
||||
} |
|
||||
shard.RUnlock() |
|
||||
wg.Done() |
|
||||
}(shard) |
|
||||
} |
|
||||
wg.Wait() |
|
||||
close(ch) |
|
||||
}() |
|
||||
|
|
||||
// Generate keys
|
|
||||
keys := make([]string, 0, count) |
|
||||
for k := range ch { |
|
||||
keys = append(keys, k) |
|
||||
} |
|
||||
return keys |
|
||||
} |
|
||||
func (m ConcurrentMap) MarshalJSON() ([]byte, error) { |
|
||||
// Create a temporary map, which will hold all item spread across shards.
|
|
||||
tmp := make(map[string]interface{}) |
|
||||
|
|
||||
// Insert items to temporary map.
|
|
||||
for item := range m.IterBuffered() { |
|
||||
tmp[item.Key] = item.Val |
|
||||
} |
|
||||
return json.Marshal(tmp) |
|
||||
} |
|
||||
|
|
||||
type UpsertCb func(exist bool, valueInMap interface{}, newValue interface{}) interface{} |
|
||||
type RemoveCb func(key string, v interface{}, exists bool) bool |
|
||||
type Tuple struct { |
|
||||
Key string |
|
||||
Val interface{} |
|
||||
} |
|
||||
|
|
||||
func snapshot(m ConcurrentMap) (chans []chan Tuple) { |
|
||||
chans = make([]chan Tuple, SHARD_COUNT) |
|
||||
wg := sync.WaitGroup{} |
|
||||
wg.Add(SHARD_COUNT) |
|
||||
// Foreach shard.
|
|
||||
for index, shard := range m { |
|
||||
go func(index int, shard *ConcurrentMapShared) { |
|
||||
// Foreach key, value pair.
|
|
||||
shard.RLock() |
|
||||
chans[index] = make(chan Tuple, len(shard.items)) |
|
||||
wg.Done() |
|
||||
for key, val := range shard.items { |
|
||||
chans[index] <- Tuple{key, val} |
|
||||
} |
|
||||
shard.RUnlock() |
|
||||
close(chans[index]) |
|
||||
}(index, shard) |
|
||||
} |
|
||||
wg.Wait() |
|
||||
return chans |
|
||||
} |
|
||||
|
|
||||
type IterCb func(key string, v interface{}) |
|
||||
|
|
||||
func fanIn(chans []chan Tuple, out chan Tuple) { |
|
||||
wg := sync.WaitGroup{} |
|
||||
wg.Add(len(chans)) |
|
||||
for _, ch := range chans { |
|
||||
go func(ch chan Tuple) { |
|
||||
for t := range ch { |
|
||||
out <- t |
|
||||
} |
|
||||
wg.Done() |
|
||||
}(ch) |
|
||||
} |
|
||||
wg.Wait() |
|
||||
close(out) |
|
||||
} |
|
||||
func fnv32(key string) uint32 { |
|
||||
hash := uint32(2166136261) |
|
||||
const prime32 = uint32(16777619) |
|
||||
for i := 0; i < len(key); i++ { |
|
||||
hash *= prime32 |
|
||||
hash ^= uint32(key[i]) |
|
||||
} |
|
||||
return hash |
|
||||
} |
|
||||
@ -1,243 +0,0 @@ |
|||||
package container |
|
||||
|
|
||||
// minCapacity is the smallest capacity that deque may have.
|
|
||||
// Must be power of 2 for bitwise modulus: x % n == x & (n - 1).
|
|
||||
const minCapacity = 16 |
|
||||
|
|
||||
// Deque represents a single instance of the deque data structure.
|
|
||||
type Deque struct { |
|
||||
buf []interface{} |
|
||||
head int |
|
||||
tail int |
|
||||
count int |
|
||||
minCap int |
|
||||
} |
|
||||
|
|
||||
// Len returns the number of elements currently stored in the queue.
|
|
||||
func (q *Deque) Len() int { |
|
||||
return q.count |
|
||||
} |
|
||||
|
|
||||
// PushBack appends an element to the back of the queue. Implements FIFO when
|
|
||||
// elements are removed with PopFront(), and LIFO when elements are removed
|
|
||||
// with PopBack().
|
|
||||
func (q *Deque) PushBack(elem interface{}) { |
|
||||
q.growIfFull() |
|
||||
|
|
||||
q.buf[q.tail] = elem |
|
||||
// Calculate new tail position.
|
|
||||
q.tail = q.next(q.tail) |
|
||||
q.count++ |
|
||||
} |
|
||||
|
|
||||
// PushFront prepends an element to the front of the queue.
|
|
||||
func (q *Deque) PushFront(elem interface{}) { |
|
||||
q.growIfFull() |
|
||||
|
|
||||
// Calculate new head position.
|
|
||||
q.head = q.prev(q.head) |
|
||||
q.buf[q.head] = elem |
|
||||
q.count++ |
|
||||
} |
|
||||
|
|
||||
// PopFront removes and returns the element from the front of the queue.
|
|
||||
// Implements FIFO when used with PushBack(). If the queue is empty, the call
|
|
||||
// panics.
|
|
||||
func (q *Deque) PopFront() interface{} { |
|
||||
if q.count <= 0 { |
|
||||
panic("deque: PopFront() called on empty queue") |
|
||||
} |
|
||||
ret := q.buf[q.head] |
|
||||
q.buf[q.head] = nil |
|
||||
// Calculate new head position.
|
|
||||
q.head = q.next(q.head) |
|
||||
q.count-- |
|
||||
|
|
||||
q.shrinkIfExcess() |
|
||||
return ret |
|
||||
} |
|
||||
|
|
||||
// PopBack removes and returns the element from the back of the queue.
|
|
||||
// Implements LIFO when used with PushBack(). If the queue is empty, the call
|
|
||||
// panics.
|
|
||||
func (q *Deque) PopBack() interface{} { |
|
||||
if q.count <= 0 { |
|
||||
panic("deque: PopBack() called on empty queue") |
|
||||
} |
|
||||
|
|
||||
// Calculate new tail position
|
|
||||
q.tail = q.prev(q.tail) |
|
||||
|
|
||||
// Remove value at tail.
|
|
||||
ret := q.buf[q.tail] |
|
||||
q.buf[q.tail] = nil |
|
||||
q.count-- |
|
||||
|
|
||||
q.shrinkIfExcess() |
|
||||
return ret |
|
||||
} |
|
||||
|
|
||||
// Front returns the element at the front of the queue. This is the element
|
|
||||
// that would be returned by PopFront(). This call panics if the queue is
|
|
||||
// empty.
|
|
||||
func (q *Deque) Front() interface{} { |
|
||||
if q.count <= 0 { |
|
||||
panic("deque: Front() called when empty") |
|
||||
} |
|
||||
return q.buf[q.head] |
|
||||
} |
|
||||
|
|
||||
// Back returns the element at the back of the queue. This is the element
|
|
||||
// that would be returned by PopBack(). This call panics if the queue is
|
|
||||
// empty.
|
|
||||
func (q *Deque) Back() interface{} { |
|
||||
if q.count <= 0 { |
|
||||
panic("deque: Back() called when empty") |
|
||||
} |
|
||||
return q.buf[q.prev(q.tail)] |
|
||||
} |
|
||||
|
|
||||
// At returns the element at index i in the queue without removing the element
|
|
||||
// from the queue. This method accepts only non-negative index values. At(0)
|
|
||||
// refers to the first element and is the same as Front(). At(Len()-1) refers
|
|
||||
// to the last element and is the same as Back(). If the index is invalid, the
|
|
||||
// call panics.
|
|
||||
//
|
|
||||
// The purpose of At is to allow Deque to serve as a more general purpose
|
|
||||
// circular buffer, where items are only added to and removed from the ends of
|
|
||||
// the deque, but may be read from any place within the deque. Consider the
|
|
||||
// case of a fixed-size circular log buffer: A new entry is pushed onto one end
|
|
||||
// and when full the oldest is popped from the other end. All the log entries
|
|
||||
// in the buffer must be readable without altering the buffer contents.
|
|
||||
func (q *Deque) At(i int) interface{} { |
|
||||
if i < 0 || i >= q.count { |
|
||||
panic("deque: At() called with index out of range") |
|
||||
} |
|
||||
// bitwise modulus
|
|
||||
return q.buf[(q.head+i)&(len(q.buf)-1)] |
|
||||
} |
|
||||
|
|
||||
// Clear removes all elements from the queue, but retains the current capacity.
|
|
||||
// This is useful when repeatedly reusing the queue at high frequency to avoid
|
|
||||
// GC during reuse. The queue will not be resized smaller as long as items are
|
|
||||
// only added. Only when items are removed is the queue subject to getting
|
|
||||
// resized smaller.
|
|
||||
func (q *Deque) Clear() { |
|
||||
// bitwise modulus
|
|
||||
modBits := len(q.buf) - 1 |
|
||||
for h := q.head; h != q.tail; h = (h + 1) & modBits { |
|
||||
q.buf[h] = nil |
|
||||
} |
|
||||
q.head = 0 |
|
||||
q.tail = 0 |
|
||||
q.count = 0 |
|
||||
} |
|
||||
|
|
||||
// Rotate rotates the deque n steps front-to-back. If n is negative, rotates
|
|
||||
// back-to-front. Having Deque provide Rotate() avoids resizing that could
|
|
||||
// happen if implementing rotation using only Pop and Push methods.
|
|
||||
func (q *Deque) Rotate(n int) { |
|
||||
if q.count <= 1 { |
|
||||
return |
|
||||
} |
|
||||
// Rotating a multiple of q.count is same as no rotation.
|
|
||||
n %= q.count |
|
||||
if n == 0 { |
|
||||
return |
|
||||
} |
|
||||
|
|
||||
modBits := len(q.buf) - 1 |
|
||||
// If no empty space in buffer, only move head and tail indexes.
|
|
||||
if q.head == q.tail { |
|
||||
// Calculate new head and tail using bitwise modulus.
|
|
||||
q.head = (q.head + n) & modBits |
|
||||
q.tail = (q.tail + n) & modBits |
|
||||
return |
|
||||
} |
|
||||
|
|
||||
if n < 0 { |
|
||||
// Rotate back to front.
|
|
||||
for ; n < 0; n++ { |
|
||||
// Calculate new head and tail using bitwise modulus.
|
|
||||
q.head = (q.head - 1) & modBits |
|
||||
q.tail = (q.tail - 1) & modBits |
|
||||
// Put tail value at head and remove value at tail.
|
|
||||
q.buf[q.head] = q.buf[q.tail] |
|
||||
q.buf[q.tail] = nil |
|
||||
} |
|
||||
return |
|
||||
} |
|
||||
|
|
||||
// Rotate front to back.
|
|
||||
for ; n > 0; n-- { |
|
||||
// Put head value at tail and remove value at head.
|
|
||||
q.buf[q.tail] = q.buf[q.head] |
|
||||
q.buf[q.head] = nil |
|
||||
// Calculate new head and tail using bitwise modulus.
|
|
||||
q.head = (q.head + 1) & modBits |
|
||||
q.tail = (q.tail + 1) & modBits |
|
||||
} |
|
||||
} |
|
||||
|
|
||||
// SetMinCapacity sets a minimum capacity of 2^minCapacityExp. If the value of
|
|
||||
// the minimum capacity is less than or equal to the minimum allowed, then
|
|
||||
// capacity is set to the minimum allowed. This may be called at anytime to
|
|
||||
// set a new minimum capacity.
|
|
||||
//
|
|
||||
// Setting a larger minimum capacity may be used to prevent resizing when the
|
|
||||
// number of stored items changes frequently across a wide range.
|
|
||||
func (q *Deque) SetMinCapacity(minCapacityExp uint) { |
|
||||
if 1<<minCapacityExp > minCapacity { |
|
||||
q.minCap = 1 << minCapacityExp |
|
||||
} else { |
|
||||
q.minCap = minCapacity |
|
||||
} |
|
||||
} |
|
||||
|
|
||||
// prev returns the previous buffer position wrapping around buffer.
|
|
||||
func (q *Deque) prev(i int) int { |
|
||||
return (i - 1) & (len(q.buf) - 1) // bitwise modulus
|
|
||||
} |
|
||||
|
|
||||
// next returns the next buffer position wrapping around buffer.
|
|
||||
func (q *Deque) next(i int) int { |
|
||||
return (i + 1) & (len(q.buf) - 1) // bitwise modulus
|
|
||||
} |
|
||||
|
|
||||
// growIfFull resizes up if the buffer is full.
|
|
||||
func (q *Deque) growIfFull() { |
|
||||
if len(q.buf) == 0 { |
|
||||
if q.minCap == 0 { |
|
||||
q.minCap = minCapacity |
|
||||
} |
|
||||
q.buf = make([]interface{}, q.minCap) |
|
||||
return |
|
||||
} |
|
||||
if q.count == len(q.buf) { |
|
||||
q.resize() |
|
||||
} |
|
||||
} |
|
||||
|
|
||||
// shrinkIfExcess resize down if the buffer 1/4 full.
|
|
||||
func (q *Deque) shrinkIfExcess() { |
|
||||
if len(q.buf) > q.minCap && (q.count<<2) == len(q.buf) { |
|
||||
q.resize() |
|
||||
} |
|
||||
} |
|
||||
|
|
||||
// resize resizes the deque to fit exactly twice its current contents. This is
|
|
||||
// used to grow the queue when it is full, and also to shrink it when it is
|
|
||||
// only a quarter full.
|
|
||||
func (q *Deque) resize() { |
|
||||
newBuf := make([]interface{}, q.count<<1) |
|
||||
if q.tail > q.head { |
|
||||
copy(newBuf, q.buf[q.head:q.tail]) |
|
||||
} else { |
|
||||
n := copy(newBuf, q.buf[q.head:]) |
|
||||
copy(newBuf[n:], q.buf[:q.tail]) |
|
||||
} |
|
||||
|
|
||||
q.head = 0 |
|
||||
q.tail = q.count |
|
||||
q.buf = newBuf |
|
||||
} |
|
||||
@ -1,51 +0,0 @@ |
|||||
package container |
|
||||
|
|
||||
// LimitedQueue 是一个具有有限容量的泛型队列
|
|
||||
type LimitedQueue[T any] struct { |
|
||||
capacity int |
|
||||
queue []T |
|
||||
} |
|
||||
|
|
||||
// NewLimitedQueue 创建一个新的 LimitedQueue
|
|
||||
func NewLimitedQueue[T any](capacity int) *LimitedQueue[T] { |
|
||||
return &LimitedQueue[T]{ |
|
||||
capacity: capacity, |
|
||||
queue: make([]T, 0, capacity), |
|
||||
} |
|
||||
} |
|
||||
|
|
||||
// Enqueue 向队列中添加一个元素
|
|
||||
func (lq *LimitedQueue[T]) Enqueue(value T) { |
|
||||
if len(lq.queue) >= lq.capacity { |
|
||||
lq.queue = lq.queue[1:] |
|
||||
} |
|
||||
lq.queue = append(lq.queue, value) |
|
||||
} |
|
||||
|
|
||||
// Dequeue 从队列中移除最早添加的元素
|
|
||||
func (lq *LimitedQueue[T]) Dequeue() *T { |
|
||||
if len(lq.queue) == 0 { |
|
||||
return nil |
|
||||
} |
|
||||
elem := lq.queue[0] |
|
||||
lq.queue = lq.queue[1:] |
|
||||
return &elem |
|
||||
} |
|
||||
|
|
||||
// Len 返回队列的长度
|
|
||||
func (lq *LimitedQueue[T]) Len() int { |
|
||||
return len(lq.queue) |
|
||||
} |
|
||||
|
|
||||
// Peek 返回队列的第一个元素但不移除它
|
|
||||
func (lq *LimitedQueue[T]) Peek() *T { |
|
||||
if len(lq.queue) == 0 { |
|
||||
return nil |
|
||||
} |
|
||||
return &lq.queue[0] |
|
||||
} |
|
||||
|
|
||||
// ToSlice 返回队列中所有元素的切片
|
|
||||
func (lq *LimitedQueue[T]) ToSlice() []T { |
|
||||
return append([]T(nil), lq.queue...) |
|
||||
} |
|
||||
@ -1,46 +0,0 @@ |
|||||
package container |
|
||||
|
|
||||
import ( |
|
||||
"container/list" |
|
||||
"fmt" |
|
||||
"sync" |
|
||||
) |
|
||||
|
|
||||
type Queue struct { |
|
||||
sync.Mutex |
|
||||
data *list.List |
|
||||
} |
|
||||
|
|
||||
func NewQueue() *Queue { |
|
||||
q := new(Queue) |
|
||||
q.data = list.New() |
|
||||
return q |
|
||||
} |
|
||||
|
|
||||
func (q *Queue) Len() int { |
|
||||
return q.data.Len() |
|
||||
} |
|
||||
|
|
||||
func (q *Queue) Push(v interface{}) { |
|
||||
defer q.Unlock() |
|
||||
q.Lock() |
|
||||
q.data.PushFront(v) |
|
||||
} |
|
||||
|
|
||||
func (q *Queue) Pop() interface{} { |
|
||||
defer q.Unlock() |
|
||||
q.Lock() |
|
||||
iter := q.data.Back() |
|
||||
if iter == nil { |
|
||||
return nil |
|
||||
} |
|
||||
v := iter.Value |
|
||||
q.data.Remove(iter) |
|
||||
return v |
|
||||
} |
|
||||
|
|
||||
func (q *Queue) Dump() { |
|
||||
for iter := q.data.Back(); iter != nil; iter = iter.Prev() { |
|
||||
fmt.Println("item:", iter.Value) |
|
||||
} |
|
||||
} |
|
||||
@ -1,112 +0,0 @@ |
|||||
package addr |
|
||||
|
|
||||
import ( |
|
||||
"fmt" |
|
||||
"net" |
|
||||
) |
|
||||
|
|
||||
var ( |
|
||||
privateBlocks []*net.IPNet |
|
||||
) |
|
||||
|
|
||||
func init() { |
|
||||
for _, b := range []string{"10.0.0.0/8", "172.16.0.0/12", "192.168.0.0/16", "100.64.0.0/10"} { |
|
||||
if _, block, err := net.ParseCIDR(b); err == nil { |
|
||||
privateBlocks = append(privateBlocks, block) |
|
||||
} |
|
||||
} |
|
||||
} |
|
||||
|
|
||||
func isPrivateIP(ipAddr string) bool { |
|
||||
ip := net.ParseIP(ipAddr) |
|
||||
for _, priv := range privateBlocks { |
|
||||
if priv.Contains(ip) { |
|
||||
return true |
|
||||
} |
|
||||
} |
|
||||
return false |
|
||||
} |
|
||||
|
|
||||
// Extract returns a real ip
|
|
||||
func Extract(addr string) (string, error) { |
|
||||
// if addr specified then its returned
|
|
||||
if len(addr) > 0 && (addr != "0.0.0.0" && addr != "[::]") { |
|
||||
return addr, nil |
|
||||
} |
|
||||
|
|
||||
addrs, err := net.InterfaceAddrs() |
|
||||
if err != nil { |
|
||||
return "", fmt.Errorf("Failed to get interface addresses! Err: %v", err) |
|
||||
} |
|
||||
|
|
||||
var ipAddr []byte |
|
||||
|
|
||||
for _, rawAddr := range addrs { |
|
||||
var ip net.IP |
|
||||
switch addr := rawAddr.(type) { |
|
||||
case *net.IPAddr: |
|
||||
ip = addr.IP |
|
||||
case *net.IPNet: |
|
||||
ip = addr.IP |
|
||||
default: |
|
||||
continue |
|
||||
} |
|
||||
|
|
||||
if ip.To4() == nil { |
|
||||
continue |
|
||||
} |
|
||||
|
|
||||
if !isPrivateIP(ip.String()) { |
|
||||
continue |
|
||||
} |
|
||||
|
|
||||
ipAddr = ip |
|
||||
break |
|
||||
} |
|
||||
|
|
||||
if ipAddr == nil { |
|
||||
return "", fmt.Errorf("No private IP address found, and explicit IP not provided") |
|
||||
} |
|
||||
|
|
||||
return net.IP(ipAddr).String(), nil |
|
||||
} |
|
||||
|
|
||||
// IPs returns all known ips
|
|
||||
func IPs() []string { |
|
||||
ifaces, err := net.Interfaces() |
|
||||
if err != nil { |
|
||||
return nil |
|
||||
} |
|
||||
|
|
||||
var ipAddrs []string |
|
||||
|
|
||||
for _, i := range ifaces { |
|
||||
addrs, err := i.Addrs() |
|
||||
if err != nil { |
|
||||
continue |
|
||||
} |
|
||||
|
|
||||
for _, addr := range addrs { |
|
||||
var ip net.IP |
|
||||
switch v := addr.(type) { |
|
||||
case *net.IPNet: |
|
||||
ip = v.IP |
|
||||
case *net.IPAddr: |
|
||||
ip = v.IP |
|
||||
} |
|
||||
|
|
||||
if ip == nil { |
|
||||
continue |
|
||||
} |
|
||||
|
|
||||
ip = ip.To4() |
|
||||
if ip == nil { |
|
||||
continue |
|
||||
} |
|
||||
|
|
||||
ipAddrs = append(ipAddrs, ip.String()) |
|
||||
} |
|
||||
} |
|
||||
|
|
||||
return ipAddrs |
|
||||
} |
|
||||
@ -1,38 +0,0 @@ |
|||||
package addr |
|
||||
|
|
||||
import ( |
|
||||
"net" |
|
||||
"testing" |
|
||||
) |
|
||||
|
|
||||
func TestExtractor(t *testing.T) { |
|
||||
testData := []struct { |
|
||||
addr string |
|
||||
expect string |
|
||||
parse bool |
|
||||
}{ |
|
||||
{"127.0.0.1", "127.0.0.1", false}, |
|
||||
{"10.0.0.1", "10.0.0.1", false}, |
|
||||
{"", "", true}, |
|
||||
{"0.0.0.0", "", true}, |
|
||||
{"[::]", "", true}, |
|
||||
} |
|
||||
|
|
||||
for _, d := range testData { |
|
||||
addr, err := Extract(d.addr) |
|
||||
if err != nil { |
|
||||
t.Errorf("Unexpected error %v", err) |
|
||||
} |
|
||||
|
|
||||
if d.parse { |
|
||||
ip := net.ParseIP(addr) |
|
||||
if ip == nil { |
|
||||
t.Error("Unexpected nil IP") |
|
||||
} |
|
||||
|
|
||||
} else if addr != d.expect { |
|
||||
t.Errorf("Expected %s got %s", d.expect, addr) |
|
||||
} |
|
||||
} |
|
||||
|
|
||||
} |
|
||||
@ -1,55 +0,0 @@ |
|||||
package ip |
|
||||
|
|
||||
import ( |
|
||||
"io/ioutil" |
|
||||
"net" |
|
||||
"net/http" |
|
||||
|
|
||||
json "github.com/json-iterator/go" |
|
||||
|
|
||||
"github.com/axgle/mahonia" |
|
||||
) |
|
||||
|
|
||||
type IPInfo struct { |
|
||||
IP string `json:"ip"` |
|
||||
Pro string `json:"pro"` |
|
||||
ProCode string `json:"proCode"` |
|
||||
City string `json:"city"` |
|
||||
CityCode string `json:"cityCode"` |
|
||||
Region string `json:"region"` |
|
||||
RegionCode string `json:"regionCode"` |
|
||||
Addr string `json:"addr"` |
|
||||
RegionNames string `json:"regionNames"` |
|
||||
Err string `json:"err"` |
|
||||
} |
|
||||
|
|
||||
// 获取以太网IP
|
|
||||
func GetEthernetInfo() *IPInfo { |
|
||||
url := "http://whois.pconline.com.cn/ipJson.jsp?json=true" |
|
||||
resp, err := http.Get(url) |
|
||||
if err != nil { |
|
||||
return nil |
|
||||
} |
|
||||
defer resp.Body.Close() |
|
||||
body, err := ioutil.ReadAll(resp.Body) |
|
||||
if err != nil { |
|
||||
return nil |
|
||||
} |
|
||||
bodystr := mahonia.NewDecoder("gbk").ConvertString(string(body)) |
|
||||
var result IPInfo |
|
||||
if err := json.Unmarshal([]byte(bodystr), &result); err != nil { |
|
||||
return nil |
|
||||
} |
|
||||
return &result |
|
||||
} |
|
||||
|
|
||||
// 获取本地Ip
|
|
||||
func GetOutboundIP() string { |
|
||||
conn, err := net.Dial("udp", "8.8.8.8:80") |
|
||||
if err != nil { |
|
||||
return "" |
|
||||
} |
|
||||
defer conn.Close() |
|
||||
localAddr := conn.LocalAddr().(*net.UDPAddr) |
|
||||
return localAddr.IP.String() |
|
||||
} |
|
||||
@ -1,83 +0,0 @@ |
|||||
package container |
|
||||
|
|
||||
import ( |
|
||||
"sync/atomic" |
|
||||
"unsafe" |
|
||||
) |
|
||||
|
|
||||
// LKQueue is a lock-free unbounded queue.
|
|
||||
type LKQueue struct { |
|
||||
head unsafe.Pointer |
|
||||
tail unsafe.Pointer |
|
||||
size int32 |
|
||||
} |
|
||||
type node struct { |
|
||||
value interface{} |
|
||||
next unsafe.Pointer |
|
||||
} |
|
||||
|
|
||||
// NewLKQueue returns an empty queue.
|
|
||||
func NewLKQueue() *LKQueue { |
|
||||
n := unsafe.Pointer(&node{}) |
|
||||
return &LKQueue{head: n, tail: n, size: 0} |
|
||||
} |
|
||||
|
|
||||
// Enqueue puts the given value v at the tail of the queue.
|
|
||||
func (q *LKQueue) Enqueue(v interface{}) { |
|
||||
n := &node{value: v} |
|
||||
for { |
|
||||
tail := load(&q.tail) |
|
||||
next := load(&tail.next) |
|
||||
if tail == load(&q.tail) { // are tail and next consistent?
|
|
||||
if next == nil { |
|
||||
if cas(&tail.next, next, n) { |
|
||||
cas(&q.tail, tail, n) // Enqueue is done. try to swing tail to the inserted node
|
|
||||
atomic.AddInt32(&q.size, 1) |
|
||||
return |
|
||||
} |
|
||||
} else { // tail was not pointing to the last node
|
|
||||
// try to swing Tail to the next node
|
|
||||
cas(&q.tail, tail, next) |
|
||||
} |
|
||||
} |
|
||||
} |
|
||||
} |
|
||||
|
|
||||
// Dequeue removes and returns the value at the head of the queue.
|
|
||||
// It returns nil if the queue is empty.
|
|
||||
func (q *LKQueue) Dequeue() interface{} { |
|
||||
for { |
|
||||
head := load(&q.head) |
|
||||
tail := load(&q.tail) |
|
||||
next := load(&head.next) |
|
||||
if head == load(&q.head) { // are head, tail, and next consistent?
|
|
||||
if head == tail { // is queue empty or tail falling behind?
|
|
||||
if next == nil { // is queue empty?
|
|
||||
return nil |
|
||||
} |
|
||||
// tail is falling behind. try to advance it
|
|
||||
cas(&q.tail, tail, next) |
|
||||
} else { |
|
||||
// read value before CAS otherwise another dequeue might free the next node
|
|
||||
v := next.value |
|
||||
if cas(&q.head, head, next) { |
|
||||
atomic.AddInt32(&q.size, -1) |
|
||||
return v // Dequeue is done. return
|
|
||||
} |
|
||||
} |
|
||||
} |
|
||||
} |
|
||||
} |
|
||||
|
|
||||
func (q *LKQueue) Size() int32 { |
|
||||
return atomic.LoadInt32(&q.size) |
|
||||
} |
|
||||
|
|
||||
func load(p *unsafe.Pointer) (n *node) { |
|
||||
return (*node)(atomic.LoadPointer(p)) |
|
||||
} |
|
||||
|
|
||||
func cas(p *unsafe.Pointer, old, new *node) (ok bool) { |
|
||||
return atomic.CompareAndSwapPointer( |
|
||||
p, unsafe.Pointer(old), unsafe.Pointer(new)) |
|
||||
} |
|
||||
@ -1,35 +0,0 @@ |
|||||
package sortslice |
|
||||
|
|
||||
//排序工具
|
|
||||
func quickSort(arr []interface{}, start, end int, compete func(a interface{}, b interface{}) int8) { |
|
||||
if start < end { |
|
||||
i, j := start, end |
|
||||
key := arr[(start+end)/2] |
|
||||
for i <= j { |
|
||||
for compete(arr[i], key) == -1 { |
|
||||
i++ |
|
||||
} |
|
||||
for compete(arr[j], key) == 1 { |
|
||||
j-- |
|
||||
} |
|
||||
if i <= j { |
|
||||
arr[i], arr[j] = arr[j], arr[i] |
|
||||
i++ |
|
||||
j-- |
|
||||
} |
|
||||
} |
|
||||
if start < j { |
|
||||
quickSort(arr, start, j, compete) |
|
||||
} |
|
||||
if end > i { |
|
||||
quickSort(arr, i, end, compete) |
|
||||
} |
|
||||
} |
|
||||
} |
|
||||
|
|
||||
func Sort(a []interface{}, compete func(a interface{}, b interface{}) int8) { |
|
||||
if len(a) < 2 { |
|
||||
return |
|
||||
} |
|
||||
quickSort(a, 0, len(a)-1, compete) |
|
||||
} |
|
||||
@ -1,27 +0,0 @@ |
|||||
package sortslice |
|
||||
|
|
||||
import ( |
|
||||
"fmt" |
|
||||
"hash/crc32" |
|
||||
"sort" |
|
||||
) |
|
||||
|
|
||||
type uInt32Slice []uint32 |
|
||||
|
|
||||
func (p uInt32Slice) Len() int { return len(p) } |
|
||||
func (p uInt32Slice) Less(i, j int) bool { return p[i] < p[j] } |
|
||||
func (p uInt32Slice) Swap(i, j int) { p[i], p[j] = p[j], p[i] } |
|
||||
|
|
||||
// Sort is a convenience method.
|
|
||||
func (p uInt32Slice) Sort() { sort.Sort(p) } |
|
||||
|
|
||||
func GetSessionId(session []uint32) (Id uint32) { |
|
||||
s := uInt32Slice(session) |
|
||||
s.Sort() |
|
||||
Session := "" |
|
||||
for _, v := range s { |
|
||||
Session = Session + fmt.Sprintf("%d", v) |
|
||||
} |
|
||||
Id = crc32.ChecksumIEEE([]byte(Session)) |
|
||||
return |
|
||||
} |
|
||||
@ -1,15 +0,0 @@ |
|||||
package version |
|
||||
|
|
||||
import ( |
|
||||
"fmt" |
|
||||
"testing" |
|
||||
"time" |
|
||||
) |
|
||||
|
|
||||
func Test_Version(t *testing.T) { |
|
||||
versionA := "1.2.3a " |
|
||||
versionB := "1.2.3b " |
|
||||
fmt.Println(CompareStrVer(versionA, versionB)) |
|
||||
time.LoadLocation("Local ") |
|
||||
fmt.Println(time.Now()) |
|
||||
} |
|
||||
@ -1,66 +0,0 @@ |
|||||
package version |
|
||||
|
|
||||
import ( |
|
||||
"strings" |
|
||||
) |
|
||||
|
|
||||
func CompareStrVer(verA, verB string) int8 { |
|
||||
verStrArrA := spliteStrByNet(verA) |
|
||||
verStrArrB := spliteStrByNet(verB) |
|
||||
lenStrA := len(verStrArrA) |
|
||||
lenStrB := len(verStrArrB) |
|
||||
if lenStrA > lenStrB { |
|
||||
return 1 |
|
||||
} |
|
||||
if lenStrA < lenStrB { |
|
||||
return -1 |
|
||||
} |
|
||||
return compareArrStrVers(verStrArrA, verStrArrB) |
|
||||
} |
|
||||
|
|
||||
// 比较版本号字符串数组
|
|
||||
func compareArrStrVers(verA, verB []string) int8 { |
|
||||
for index, _ := range verA { |
|
||||
littleResult := compareLittleVer(verA[index], verB[index]) |
|
||||
if littleResult != 0 { |
|
||||
return littleResult |
|
||||
} |
|
||||
} |
|
||||
return 0 |
|
||||
} |
|
||||
|
|
||||
//
|
|
||||
// 比较小版本号字符串
|
|
||||
//
|
|
||||
func compareLittleVer(verA, verB string) int8 { |
|
||||
bytesA := []byte(verA) |
|
||||
bytesB := []byte(verB) |
|
||||
lenA := len(bytesA) |
|
||||
lenB := len(bytesB) |
|
||||
if lenA > lenB { |
|
||||
return 1 |
|
||||
} |
|
||||
if lenA < lenB { |
|
||||
return -1 |
|
||||
} |
|
||||
//如果长度相等则按byte位进行比较
|
|
||||
return compareByBytes(bytesA, bytesB) |
|
||||
} |
|
||||
|
|
||||
// 按byte位进行比较小版本号
|
|
||||
func compareByBytes(verA, verB []byte) int8 { |
|
||||
for index, _ := range verA { |
|
||||
if verA[index] > verB[index] { |
|
||||
return 1 |
|
||||
} |
|
||||
if verA[index] < verB[index] { |
|
||||
return -1 |
|
||||
} |
|
||||
} |
|
||||
return 0 |
|
||||
} |
|
||||
|
|
||||
// 按“.”分割版本号为小版本号的字符串数组
|
|
||||
func spliteStrByNet(strV string) []string { |
|
||||
return strings.Split(strV, ". ") |
|
||||
} |
|
||||
@ -1,13 +0,0 @@ |
|||||
package base64 |
|
||||
|
|
||||
import "encoding/base64" |
|
||||
|
|
||||
///编码
|
|
||||
func EncodeToString(src []byte) string { |
|
||||
return base64.StdEncoding.EncodeToString(src) |
|
||||
} |
|
||||
|
|
||||
///解码
|
|
||||
func DecodeString(s string) ([]byte, error) { |
|
||||
return base64.StdEncoding.DecodeString(s) |
|
||||
} |
|
||||
@ -1,116 +0,0 @@ |
|||||
package crypto |
|
||||
|
|
||||
/* |
|
||||
国密算法库 github.com/tjfoc/gmsm/ 封装 |
|
||||
*/ |
|
||||
|
|
||||
import ( |
|
||||
"bytes" |
|
||||
"crypto/rand" |
|
||||
|
|
||||
"github.com/tjfoc/gmsm/sm2" |
|
||||
"github.com/tjfoc/gmsm/sm3" |
|
||||
"github.com/tjfoc/gmsm/sm4" |
|
||||
"github.com/tjfoc/gmsm/x509" |
|
||||
) |
|
||||
|
|
||||
///生成密钥对
|
|
||||
func GenerateKey(key string) (*sm2.PrivateKey, error) { |
|
||||
return sm2.GenerateKey(bytes.NewBufferString(key)) |
|
||||
} |
|
||||
|
|
||||
///读取私钥
|
|
||||
func ReadPrivateKeyFromPem(privateKeyPem []byte, pwd []byte) (*sm2.PrivateKey, error) { |
|
||||
return x509.ReadPrivateKeyFromPem(privateKeyPem, pwd) |
|
||||
} |
|
||||
|
|
||||
///读取私钥
|
|
||||
func ReadPublicKeyFromPem(privateKeyPem []byte) (*sm2.PublicKey, error) { |
|
||||
return x509.ReadPublicKeyFromPem(privateKeyPem) |
|
||||
} |
|
||||
|
|
||||
///国密—SM2-加密
|
|
||||
func GM_SM2_Encry(origData []byte, pub *sm2.PublicKey) (ciphertext []byte, err error) { |
|
||||
ciphertext, err = pub.EncryptAsn1(origData, rand.Reader) |
|
||||
return |
|
||||
} |
|
||||
|
|
||||
///国密—SM2-解密
|
|
||||
func GM_SM2_Decry(ciphertext []byte, priv *sm2.PrivateKey) (origData []byte, err error) { |
|
||||
origData, err = priv.DecryptAsn1(ciphertext) |
|
||||
return |
|
||||
} |
|
||||
|
|
||||
///国密—SM2-签名
|
|
||||
func GM_SM2_Sign(origData, key string) (sign string, err error) { |
|
||||
var ( |
|
||||
priv *sm2.PrivateKey |
|
||||
msg []byte |
|
||||
signdata []byte |
|
||||
) |
|
||||
if priv, err = sm2.GenerateKey(bytes.NewBufferString(key)); err != nil { |
|
||||
return |
|
||||
} |
|
||||
msg = []byte(origData) |
|
||||
if signdata, err = priv.Sign(rand.Reader, msg, nil); err != nil { |
|
||||
return |
|
||||
} |
|
||||
sign = string(signdata) |
|
||||
return |
|
||||
} |
|
||||
|
|
||||
///国密—SM2-验签
|
|
||||
func GM_SM2_Verify(origData, sign, key string) (isok bool, err error) { |
|
||||
var ( |
|
||||
priv *sm2.PrivateKey |
|
||||
pub *sm2.PublicKey |
|
||||
) |
|
||||
if priv, err = sm2.GenerateKey(bytes.NewBufferString(key)); err != nil { |
|
||||
return |
|
||||
} |
|
||||
pub = &priv.PublicKey |
|
||||
isok = pub.Verify([]byte(origData), []byte(sign)) |
|
||||
return |
|
||||
} |
|
||||
|
|
||||
///国密—SM3-哈希
|
|
||||
func GM_SM3_Hash(origData string) (hash string) { |
|
||||
h := sm3.New() |
|
||||
h.Write([]byte(origData)) |
|
||||
hash = string(h.Sum(nil)) |
|
||||
return |
|
||||
} |
|
||||
|
|
||||
///国密-SM4-sm4Ecb模式pksc7填充加密
|
|
||||
func GM_SM4_Ecb(origData string, key string) (ciphertext string, err error) { |
|
||||
var ( |
|
||||
iv []byte |
|
||||
ecbdata []byte |
|
||||
) |
|
||||
iv = []byte("0000000000000000") |
|
||||
if err = sm4.SetIV(iv); err != nil { //设置SM4算法实现的IV值,不设置则使用默认值
|
|
||||
return |
|
||||
} |
|
||||
if ecbdata, err = sm4.Sm4Ecb([]byte(key), []byte(origData), true); err != nil { //sm4Ecb模式pksc7填充加密
|
|
||||
return |
|
||||
} |
|
||||
ciphertext = string(ecbdata) |
|
||||
return |
|
||||
} |
|
||||
|
|
||||
///国密-SM4-sm4Ecb模式pksc7填充加密
|
|
||||
func GM_SM4_Dec(ciphertext string, key string) (origData string, err error) { |
|
||||
var ( |
|
||||
iv []byte |
|
||||
ecbdata []byte |
|
||||
) |
|
||||
iv = []byte("0000000000000000") |
|
||||
if err = sm4.SetIV(iv); err != nil { //设置SM4算法实现的IV值,不设置则使用默认值
|
|
||||
return |
|
||||
} |
|
||||
if ecbdata, err = sm4.Sm4Ecb([]byte(key), []byte(ciphertext), false); err != nil { //sm4Ecb模式pksc7填充加密
|
|
||||
return |
|
||||
} |
|
||||
origData = string(ecbdata) |
|
||||
return |
|
||||
} |
|
||||
@ -1,87 +0,0 @@ |
|||||
package crypto |
|
||||
|
|
||||
import ( |
|
||||
"crypto/rand" |
|
||||
"fmt" |
|
||||
"testing" |
|
||||
|
|
||||
"github.com/tjfoc/gmsm/sm2" |
|
||||
"github.com/tjfoc/gmsm/x509" |
|
||||
) |
|
||||
|
|
||||
func Test_GM_SM2_Encry(t *testing.T) { |
|
||||
priv, err := sm2.GenerateKey(nil) // 生成密钥对
|
|
||||
if err != nil { |
|
||||
t.Fatal(err) |
|
||||
} |
|
||||
privPem, err := x509.WritePrivateKeyToPem(priv, nil) // 生成密钥文件
|
|
||||
if err != nil { |
|
||||
t.Fatal(err) |
|
||||
} else { |
|
||||
fmt.Printf("privPem:%s\n", string(privPem)) |
|
||||
} |
|
||||
pubKey, _ := priv.Public().(*sm2.PublicKey) |
|
||||
pubkeyPem, err := x509.WritePublicKeyToPem(pubKey) // 生成公钥文件
|
|
||||
_, err = x509.ReadPrivateKeyFromPem(privPem, nil) // 读取密钥
|
|
||||
if err != nil { |
|
||||
t.Fatal(err) |
|
||||
} else { |
|
||||
fmt.Printf("pubkeyPem:%s\n", string(pubkeyPem)) |
|
||||
} |
|
||||
pubKey, err = x509.ReadPublicKeyFromPem(pubkeyPem) // 读取公钥
|
|
||||
if err != nil { |
|
||||
t.Fatal(err) |
|
||||
} |
|
||||
cipher, err := pubKey.EncryptAsn1([]byte( |
|
||||
`liwei1dao |
|
||||
liwei2dao |
|
||||
liwei3dao |
|
||||
liwei4dao`), rand.Reader) |
|
||||
if err != nil { |
|
||||
t.Fatal(err) |
|
||||
} else { |
|
||||
fmt.Printf("cipher:%s\n", string(cipher)) |
|
||||
} |
|
||||
|
|
||||
origData, err := priv.DecryptAsn1(cipher) |
|
||||
if err != nil { |
|
||||
t.Fatal(err) |
|
||||
} else { |
|
||||
fmt.Printf("origData:%s\n", string(origData)) |
|
||||
} |
|
||||
} |
|
||||
|
|
||||
func Test_GM_SM2_Decry(t *testing.T) { |
|
||||
|
|
||||
} |
|
||||
|
|
||||
func Test_GM_SM2_Sign(t *testing.T) { |
|
||||
origData, err := GM_SM2_Sign("token", "123456781234567812345678") |
|
||||
fmt.Printf("origData:%s err:%v", origData, err) |
|
||||
} |
|
||||
|
|
||||
func Test_GM_SM2_Verify(t *testing.T) { |
|
||||
isok, err := GM_SM2_Verify("token", "123456781234567812345678", "") |
|
||||
fmt.Printf("isok:%v rr:%v", isok, err) |
|
||||
} |
|
||||
|
|
||||
func Test_GM_SM3_Hash(t *testing.T) { |
|
||||
hash := GM_SM3_Hash("123456781234567812345678") |
|
||||
fmt.Printf("hash:%v", hash) |
|
||||
} |
|
||||
|
|
||||
func Test_GM_SM4_Ecb(t *testing.T) { |
|
||||
ciphertext, err := GM_SM4_Ecb("123456781234567812345678", "") |
|
||||
fmt.Printf("ciphertext:%v err:%v", ciphertext, err) |
|
||||
} |
|
||||
|
|
||||
func Test_GM_SM4_Dec(t *testing.T) { |
|
||||
ciphertext, err := GM_SM4_Dec("123456781234567812345678", "") |
|
||||
fmt.Printf("ciphertext:%v err:%v", ciphertext, err) |
|
||||
} |
|
||||
|
|
||||
func Test_GM_SM2_Hex(t *testing.T) { |
|
||||
// if priv, err := x509.ReadPrivateKeyFromHex(""); err != nil {
|
|
||||
|
|
||||
// }
|
|
||||
} |
|
||||
@ -1,49 +0,0 @@ |
|||||
package gm_java |
|
||||
|
|
||||
/* |
|
||||
国密算法库 github.com/ZZMarquis/gm/sm2 封装 |
|
||||
主要处理Java 版本国密算法 |
|
||||
*/ |
|
||||
|
|
||||
import ( |
|
||||
"encoding/base64" |
|
||||
"fmt" |
|
||||
|
|
||||
"github.com/ZZMarquis/gm/sm2" |
|
||||
) |
|
||||
|
|
||||
//样例 base64Key:FXlrn8jX61JDcBtOOh59yy/sM2r1hBT5XODayZKRDVE=
|
|
||||
func ReadPrivateKeyFormBase64(base64Key string) (privKey *sm2.PrivateKey, err error) { |
|
||||
var ( |
|
||||
key []byte |
|
||||
) |
|
||||
if key, err = base64.StdEncoding.DecodeString(base64Key); err != nil { |
|
||||
err = fmt.Errorf("ReadPrivateKeyFormBase64 Base64 DecodeString err:%v", err) |
|
||||
return |
|
||||
} |
|
||||
privKey, err = sm2.RawBytesToPrivateKey(key) |
|
||||
return |
|
||||
} |
|
||||
|
|
||||
//样例 base64Key:BHa3F+W4YhmWoqa7glAURrU7vijUSNtg+9ZnQREuq/8+6MsGAc7
|
|
||||
func ReadPublicKeyFormBase64(base64Key string) (pubKey *sm2.PublicKey, err error) { |
|
||||
var ( |
|
||||
key []byte |
|
||||
) |
|
||||
if key, err = base64.StdEncoding.DecodeString(base64Key); err != nil { |
|
||||
err = fmt.Errorf("ReadPublicKeyFormBase64 Base64 DecodeString err:%v", err) |
|
||||
return |
|
||||
} |
|
||||
pubKey, err = sm2.RawBytesToPublicKey(key) |
|
||||
return |
|
||||
} |
|
||||
|
|
||||
///SM2 加密
|
|
||||
func SM2_Encrypt(pubKey *sm2.PublicKey, in []byte) ([]byte, error) { |
|
||||
return sm2.Encrypt(pubKey, in, sm2.C1C3C2) |
|
||||
} |
|
||||
|
|
||||
///SM2 解密
|
|
||||
func SM2_Decrypt(privKey *sm2.PrivateKey, in []byte) ([]byte, error) { |
|
||||
return sm2.Decrypt(privKey, in, sm2.C1C3C2) |
|
||||
} |
|
||||
@ -1,21 +0,0 @@ |
|||||
package gm_java |
|
||||
|
|
||||
import ( |
|
||||
"encoding/base64" |
|
||||
"fmt" |
|
||||
"io/ioutil" |
|
||||
"testing" |
|
||||
|
|
||||
"github.com/ZZMarquis/gm/sm2" |
|
||||
) |
|
||||
|
|
||||
func Test_RSA_Decrypt(t *testing.T) { |
|
||||
data, _ := base64.StdEncoding.DecodeString("FXlrn8jX61JDcBtOOh59yy/sM2r1hBT5XODayZKRDVE=") |
|
||||
privKey, err := sm2.RawBytesToPrivateKey(data) |
|
||||
fmt.Printf("privKey:%v err:%v\n", privKey, err) |
|
||||
content, err := ioutil.ReadFile("encode.txt") |
|
||||
fmt.Printf("content:%s err:%v\n", string(content), err) |
|
||||
mdata, _ := base64.StdEncoding.DecodeString(string(content)) |
|
||||
orgdata, err := sm2.Decrypt(privKey, mdata, sm2.C1C3C2) |
|
||||
fmt.Printf("orgdata:%s err:%v\n", string(orgdata), err) |
|
||||
} |
|
||||
@ -1,133 +0,0 @@ |
|||||
package sra |
|
||||
|
|
||||
/* |
|
||||
RSA 非对称加密算法封装 |
|
||||
*/ |
|
||||
|
|
||||
import ( |
|
||||
"crypto" |
|
||||
"crypto/rand" |
|
||||
"crypto/rsa" |
|
||||
"crypto/sha256" |
|
||||
"crypto/x509" |
|
||||
"encoding/pem" |
|
||||
"errors" |
|
||||
"fmt" |
|
||||
) |
|
||||
|
|
||||
//RSA公钥私钥产生
|
|
||||
func GenRsaKey(bits int) (prvkey, pubkey []byte, err error) { |
|
||||
// 生成私钥文件
|
|
||||
privateKey, err := rsa.GenerateKey(rand.Reader, bits) |
|
||||
if err != nil { |
|
||||
return |
|
||||
} |
|
||||
derStream := x509.MarshalPKCS1PrivateKey(privateKey) |
|
||||
block := &pem.Block{ |
|
||||
Type: "RSA PRIVATE KEY", |
|
||||
Bytes: derStream, |
|
||||
} |
|
||||
prvkey = pem.EncodeToMemory(block) |
|
||||
publicKey := &privateKey.PublicKey |
|
||||
derPkix, err := x509.MarshalPKIXPublicKey(publicKey) |
|
||||
if err != nil { |
|
||||
return |
|
||||
} |
|
||||
block = &pem.Block{ |
|
||||
Type: "PUBLIC KEY", |
|
||||
Bytes: derPkix, |
|
||||
} |
|
||||
pubkey = pem.EncodeToMemory(block) |
|
||||
return |
|
||||
} |
|
||||
|
|
||||
//签名
|
|
||||
func RsaSignWithSha256(data []byte, keyBytes []byte) (sgin []byte, err error) { |
|
||||
var ( |
|
||||
privateKey *rsa.PrivateKey |
|
||||
) |
|
||||
h := sha256.New() |
|
||||
h.Write(data) |
|
||||
hashed := h.Sum(nil) |
|
||||
block, _ := pem.Decode(keyBytes) |
|
||||
if block == nil { |
|
||||
err = errors.New("private key error") |
|
||||
return |
|
||||
} |
|
||||
privateKey, err = x509.ParsePKCS1PrivateKey(block.Bytes) |
|
||||
if err != nil { |
|
||||
err = fmt.Errorf("ParsePKCS8PrivateKey err:%v", err) |
|
||||
return |
|
||||
} |
|
||||
|
|
||||
sgin, err = rsa.SignPKCS1v15(rand.Reader, privateKey, crypto.SHA256, hashed) |
|
||||
if err != nil { |
|
||||
err = fmt.Errorf("Error from signing err:%v", err) |
|
||||
return |
|
||||
} |
|
||||
return |
|
||||
} |
|
||||
|
|
||||
//验证
|
|
||||
func RsaVerySignWithSha256(data, signData, keyBytes []byte) (issucc bool, err error) { |
|
||||
var ( |
|
||||
pubKey interface{} |
|
||||
) |
|
||||
block, _ := pem.Decode(keyBytes) |
|
||||
if block == nil { |
|
||||
err = errors.New("public key error") |
|
||||
return |
|
||||
} |
|
||||
pubKey, err = x509.ParsePKIXPublicKey(block.Bytes) |
|
||||
if err != nil { |
|
||||
return |
|
||||
} |
|
||||
hashed := sha256.Sum256(data) |
|
||||
err = rsa.VerifyPKCS1v15(pubKey.(*rsa.PublicKey), crypto.SHA256, hashed[:], signData) |
|
||||
if err != nil { |
|
||||
return |
|
||||
} |
|
||||
return true, nil |
|
||||
} |
|
||||
|
|
||||
// 公钥加密
|
|
||||
func RsaEncrypt(data, keyBytes []byte) ([]byte, error) { |
|
||||
//解密pem格式的公钥
|
|
||||
block, _ := pem.Decode(keyBytes) |
|
||||
if block == nil { |
|
||||
return nil, errors.New("public key error") |
|
||||
} |
|
||||
// 解析公钥
|
|
||||
pubInterface, err := x509.ParsePKIXPublicKey(block.Bytes) |
|
||||
if err != nil { |
|
||||
return nil, err |
|
||||
} |
|
||||
// 类型断言
|
|
||||
pub := pubInterface.(*rsa.PublicKey) |
|
||||
//加密
|
|
||||
ciphertext, err := rsa.EncryptPKCS1v15(rand.Reader, pub, data) |
|
||||
if err != nil { |
|
||||
return nil, err |
|
||||
} |
|
||||
return ciphertext, nil |
|
||||
} |
|
||||
|
|
||||
// 私钥解密
|
|
||||
func RsaDecrypt(ciphertext, keyBytes []byte) ([]byte, error) { |
|
||||
//获取私钥
|
|
||||
block, _ := pem.Decode(keyBytes) |
|
||||
if block == nil { |
|
||||
return nil, errors.New("private key error!") |
|
||||
} |
|
||||
//解析PKCS1格式的私钥
|
|
||||
priv, err := x509.ParsePKCS1PrivateKey(block.Bytes) |
|
||||
if err != nil { |
|
||||
return nil, err |
|
||||
} |
|
||||
// 解密
|
|
||||
data, err := rsa.DecryptPKCS1v15(rand.Reader, priv, ciphertext) |
|
||||
if err != nil { |
|
||||
return nil, err |
|
||||
} |
|
||||
return data, nil |
|
||||
} |
|
||||
@ -1,31 +0,0 @@ |
|||||
package sra |
|
||||
|
|
||||
import ( |
|
||||
"encoding/hex" |
|
||||
"fmt" |
|
||||
"testing" |
|
||||
) |
|
||||
|
|
||||
func Test_RSA(t *testing.T) { |
|
||||
//rsa 密钥文件产生
|
|
||||
fmt.Println("-------------------------------获取RSA公私钥-----------------------------------------") |
|
||||
prvKey, pubKey, _ := GenRsaKey(1204) |
|
||||
fmt.Println(string(prvKey)) |
|
||||
fmt.Println(string(pubKey)) |
|
||||
|
|
||||
fmt.Println("-------------------------------进行签名与验证操作-----------------------------------------") |
|
||||
var data = "卧了个槽,这么神奇的吗??!!! ԅ(¯﹃¯ԅ) !!!!!!)" |
|
||||
fmt.Println("对消息进行签名操作...") |
|
||||
signData, _ := RsaSignWithSha256([]byte(data), prvKey) |
|
||||
fmt.Println("消息的签名信息: ", hex.EncodeToString(signData)) |
|
||||
fmt.Println("\n对签名信息进行验证...") |
|
||||
if ok, _ := RsaVerySignWithSha256([]byte(data), signData, pubKey); ok { |
|
||||
fmt.Println("签名信息验证成功,确定是正确私钥签名!!") |
|
||||
} |
|
||||
|
|
||||
fmt.Println("-------------------------------进行加密解密操作-----------------------------------------") |
|
||||
ciphertext, _ := RsaEncrypt([]byte(data), pubKey) |
|
||||
fmt.Println("公钥加密后的数据:", hex.EncodeToString(ciphertext)) |
|
||||
sourceData, _ := RsaDecrypt(ciphertext, prvKey) |
|
||||
fmt.Println("私钥解密后的数据:", string(sourceData)) |
|
||||
} |
|
||||
@ -1,113 +0,0 @@ |
|||||
package mcp |
|
||||
|
|
||||
import ( |
|
||||
"context" |
|
||||
"yunyan/lego/core" |
|
||||
"yunyan/lego/core/cbase" |
|
||||
"yunyan/lego/sys/log" |
|
||||
"yunyan/sys/juhe" |
|
||||
"fmt" |
|
||||
|
|
||||
"github.com/mark3labs/mcp-go/mcp" |
|
||||
) |
|
||||
|
|
||||
// 股票信息卡片
|
|
||||
type FinanceCard struct { |
|
||||
CompanyName string `json:"companyName"` |
|
||||
CurrPrice string `json:"currPrice"` |
|
||||
OpenPrice string `json:"openPrice"` |
|
||||
TodayLow string `json:"todayLow"` |
|
||||
TodayHigh string `json:"todayHigh"` |
|
||||
Volume string `json:"volume"` |
|
||||
} |
|
||||
|
|
||||
// 股票查询
|
|
||||
type tool_finance struct { |
|
||||
cbase.ModuleCompBase |
|
||||
module *Mcp |
|
||||
} |
|
||||
|
|
||||
// 组件初始化接口
|
|
||||
func (this *tool_finance) Init(service core.IService, module core.IModule, comp core.IModuleComp, opt core.IModuleOptions) (err error) { |
|
||||
this.ModuleCompBase.Init(service, module, comp, opt) |
|
||||
this.module = module.(*Mcp) |
|
||||
return |
|
||||
} |
|
||||
func (this *tool_finance) Start() (err error) { |
|
||||
err = this.ModuleCompBase.Start() |
|
||||
if this.module.AddTool(ToolGroup_CHINA, this.Tool(), this.Handl) { |
|
||||
if err := juhe.OnInit(this.module.service.GetSettings().Sys["juhe"]); err != nil { |
|
||||
panic(fmt.Sprintf("init sys.juhe err: %s", err.Error())) |
|
||||
} else { |
|
||||
log.Infof("init sys.juhe success!") |
|
||||
} |
|
||||
} |
|
||||
return |
|
||||
} |
|
||||
func (this *tool_finance) Tool() (tool mcp.Tool) { |
|
||||
return mcp.NewTool("finance_query", |
|
||||
mcp.WithDescription("根据提供的股票代码查询股票信息"), |
|
||||
mcp.WithString("code", |
|
||||
mcp.Description("股票代码"), |
|
||||
mcp.Required(), |
|
||||
), |
|
||||
) |
|
||||
} |
|
||||
|
|
||||
// Handl 工具处理接口
|
|
||||
// 参数:
|
|
||||
// - ctx: 上下文
|
|
||||
// - request: MCP 工具调用请求
|
|
||||
//
|
|
||||
// 返回值:
|
|
||||
// - result: MCP 工具调用结果
|
|
||||
// - err: 处理过程中的错误
|
|
||||
//
|
|
||||
// 异常:
|
|
||||
// - 无(不主动 panic)
|
|
||||
func (this *tool_finance) Handl(ctx context.Context, request mcp.CallToolRequest) (result *mcp.CallToolResult, err error) { |
|
||||
var ( |
|
||||
code string |
|
||||
data *juhe.StockResponse |
|
||||
card *FinanceCard |
|
||||
context string |
|
||||
) |
|
||||
if code, err = request.RequireString("code"); err != nil { |
|
||||
return |
|
||||
} |
|
||||
if data, err = juhe.Financebygid(code); err != nil { |
|
||||
context = "获取天气信息失败" |
|
||||
} else { |
|
||||
context = fmt.Sprintf("当前%s股票情况:今日开盘价:%s,昨日收盘价:%s,当前价格:%s,今日最高价:%s,今日最低价:%s,成交量:%s,成交金额:%s", |
|
||||
data.FinanceResult[0].Data.Name, |
|
||||
data.FinanceResult[0].Data.TodayStartPri, |
|
||||
data.FinanceResult[0].Data.YestodEndPri, |
|
||||
data.FinanceResult[0].Data.NowPri, |
|
||||
data.FinanceResult[0].Data.TodayMax, |
|
||||
data.FinanceResult[0].Data.TodayMin, |
|
||||
data.FinanceResult[0].Data.TraNumber, |
|
||||
data.FinanceResult[0].Data.TraAmount, |
|
||||
) |
|
||||
card = &FinanceCard{ |
|
||||
CompanyName: code, |
|
||||
CurrPrice: data.FinanceResult[0].Data.NowPri, |
|
||||
OpenPrice: data.FinanceResult[0].Data.TodayStartPri, |
|
||||
TodayLow: data.FinanceResult[0].Data.TodayMin, |
|
||||
TodayHigh: data.FinanceResult[0].Data.TodayMax, |
|
||||
Volume: data.FinanceResult[0].Data.TraNumber, |
|
||||
} |
|
||||
} |
|
||||
|
|
||||
result = &mcp.CallToolResult{ |
|
||||
Content: []mcp.Content{ |
|
||||
mcp.TextContent{ |
|
||||
Type: "text", |
|
||||
Text: context, |
|
||||
}, |
|
||||
}, |
|
||||
} |
|
||||
result.Meta = map[string]interface{}{ |
|
||||
"tool_finance": card, |
|
||||
} |
|
||||
return |
|
||||
} |
|
||||
@ -1,122 +0,0 @@ |
|||||
package mcp |
|
||||
|
|
||||
import ( |
|
||||
"context" |
|
||||
"yunyan/lego/core" |
|
||||
"yunyan/lego/core/cbase" |
|
||||
"yunyan/lego/sys/log" |
|
||||
lgspotify "yunyan/sys/spotify" |
|
||||
"yunyan/utils" |
|
||||
"fmt" |
|
||||
|
|
||||
"github.com/mark3labs/mcp-go/mcp" |
|
||||
"github.com/zmb3/spotify/v2" |
|
||||
) |
|
||||
|
|
||||
// 音乐卡片 结构统一化输出
|
|
||||
type SpotifyMusicCard struct { |
|
||||
Name string `json:"name"` |
|
||||
Images string `json:"image"` |
|
||||
Url string `json:"url"` |
|
||||
Sgener string `json:"sgener"` |
|
||||
} |
|
||||
|
|
||||
// 天气查询
|
|
||||
type tool_spotify_music_play struct { |
|
||||
cbase.ModuleCompBase |
|
||||
module *Mcp |
|
||||
} |
|
||||
|
|
||||
// 组件初始化接口
|
|
||||
func (this *tool_spotify_music_play) Init(service core.IService, module core.IModule, comp core.IModuleComp, opt core.IModuleOptions) (err error) { |
|
||||
this.ModuleCompBase.Init(service, module, comp, opt) |
|
||||
this.module = module.(*Mcp) |
|
||||
|
|
||||
return |
|
||||
} |
|
||||
|
|
||||
func (this *tool_spotify_music_play) Start() (err error) { |
|
||||
err = this.ModuleCompBase.Start() |
|
||||
if this.module.AddTool(ToolGroup_GLOBAL, this.Tool(), this.Handl) { |
|
||||
if err := lgspotify.OnInit(this.module.service.GetSettings().Sys["spotify"]); err != nil { |
|
||||
panic(fmt.Sprintf("init sys.spotify err: %s", err.Error())) |
|
||||
} else { |
|
||||
log.Infof("init sys.spotify success!") |
|
||||
} |
|
||||
} |
|
||||
return |
|
||||
} |
|
||||
|
|
||||
func (this *tool_spotify_music_play) Tool() (tool mcp.Tool) { |
|
||||
return mcp.NewTool("spotify_music_play", |
|
||||
mcp.WithDescription("为用户提供音乐播放服务"), |
|
||||
mcp.WithString("query", |
|
||||
mcp.Description("Music keyword information"), |
|
||||
mcp.Required(), |
|
||||
), |
|
||||
mcp.WithString("offset", |
|
||||
mcp.Description("music offset (max 9, default 0)"), |
|
||||
mcp.DefaultNumber(0), |
|
||||
), |
|
||||
) |
|
||||
} |
|
||||
|
|
||||
// Handl 工具处理接口
|
|
||||
// 参数:
|
|
||||
// - ctx: 上下文
|
|
||||
// - request: MCP 工具调用请求
|
|
||||
//
|
|
||||
// 返回值:
|
|
||||
// - result: MCP 工具调用结果
|
|
||||
// - err: 处理过程中的错误
|
|
||||
//
|
|
||||
// 异常:
|
|
||||
// - 无(不主动 panic)
|
|
||||
func (this *tool_spotify_music_play) Handl(ctx context.Context, request mcp.CallToolRequest) (result *mcp.CallToolResult, err error) { |
|
||||
var ( |
|
||||
offset int |
|
||||
query string |
|
||||
data []spotify.FullTrack |
|
||||
card *SpotifyMusicCard |
|
||||
context string |
|
||||
) |
|
||||
offset = request.GetInt("offset", 0) |
|
||||
if query, err = request.RequireString("query"); err != nil { |
|
||||
return |
|
||||
} |
|
||||
data, err = lgspotify.SearchMusic(ctx, query, 1, offset) |
|
||||
|
|
||||
if err != nil { |
|
||||
context = "播放音乐失败" |
|
||||
log.Errorf("获取天气信息失败:%v", err) |
|
||||
return |
|
||||
} else { |
|
||||
card = &SpotifyMusicCard{ |
|
||||
Name: data[0].Name, |
|
||||
Images: data[0].Album.Images[0].URL, |
|
||||
Url: string(data[0].URI), |
|
||||
Sgener: data[0].Album.Name, |
|
||||
} |
|
||||
context = utils.ToString(map[string]interface{}{ |
|
||||
"card_music": card, |
|
||||
"result": fmt.Sprintf("正在为你播放音乐:%s", data[0].Name), |
|
||||
}) |
|
||||
log.Debug("播放音乐成功", |
|
||||
log.Field{Key: "query", Value: query}, |
|
||||
log.Field{Key: "card", Value: card}, |
|
||||
) |
|
||||
} |
|
||||
|
|
||||
result = &mcp.CallToolResult{ |
|
||||
Content: []mcp.Content{ |
|
||||
mcp.TextContent{ |
|
||||
Type: "text", |
|
||||
Text: context, |
|
||||
}, |
|
||||
}, |
|
||||
} |
|
||||
// result.Meta = map[string]interface{}{
|
|
||||
// "card_spotify": card,
|
|
||||
// }
|
|
||||
return |
|
||||
} |
|
||||
@ -1,134 +0,0 @@ |
|||||
package mcp |
|
||||
|
|
||||
import ( |
|
||||
"context" |
|
||||
"yunyan/lego/core" |
|
||||
"yunyan/lego/core/cbase" |
|
||||
"yunyan/lego/sys/log" |
|
||||
lgspotify "yunyan/sys/spotify" |
|
||||
"yunyan/utils" |
|
||||
"fmt" |
|
||||
|
|
||||
"github.com/mark3labs/mcp-go/mcp" |
|
||||
"github.com/zmb3/spotify/v2" |
|
||||
) |
|
||||
|
|
||||
// 音乐卡片 结构统一化输出
|
|
||||
type SpotifyMusicListCard struct { |
|
||||
Musics []SpotifyMusicCard `json:"musics"` |
|
||||
} |
|
||||
|
|
||||
// 天气查询
|
|
||||
type tool_spotify_music_playlist struct { |
|
||||
cbase.ModuleCompBase |
|
||||
module *Mcp |
|
||||
} |
|
||||
|
|
||||
// 组件初始化接口
|
|
||||
func (this *tool_spotify_music_playlist) Init(service core.IService, module core.IModule, comp core.IModuleComp, opt core.IModuleOptions) (err error) { |
|
||||
this.ModuleCompBase.Init(service, module, comp, opt) |
|
||||
this.module = module.(*Mcp) |
|
||||
|
|
||||
return |
|
||||
} |
|
||||
|
|
||||
func (this *tool_spotify_music_playlist) Start() (err error) { |
|
||||
err = this.ModuleCompBase.Start() |
|
||||
if this.module.AddTool(ToolGroup_GLOBAL, this.Tool(), this.Handl) { |
|
||||
if err := lgspotify.OnInit(this.module.service.GetSettings().Sys["spotify"]); err != nil { |
|
||||
panic(fmt.Sprintf("init sys.spotify err: %s", err.Error())) |
|
||||
} else { |
|
||||
log.Infof("init sys.spotify success!") |
|
||||
} |
|
||||
} |
|
||||
return |
|
||||
} |
|
||||
|
|
||||
func (this *tool_spotify_music_playlist) Tool() (tool mcp.Tool) { |
|
||||
return mcp.NewTool("spotify_music_play", |
|
||||
mcp.WithDescription("Recommend the desired song list to users to realize the music playing experience"), |
|
||||
mcp.WithString("query", |
|
||||
mcp.Description("Music keyword information"), |
|
||||
mcp.Required(), |
|
||||
), |
|
||||
mcp.WithNumber("count", |
|
||||
mcp.Description("music of results (1-20, default 10)"), |
|
||||
mcp.DefaultNumber(10), |
|
||||
mcp.Min(1), |
|
||||
mcp.Max(20), |
|
||||
), |
|
||||
mcp.WithString("offset", |
|
||||
mcp.Description("music offset (max 9, default 0)"), |
|
||||
mcp.DefaultNumber(0), |
|
||||
), |
|
||||
) |
|
||||
} |
|
||||
|
|
||||
// Handl 工具处理接口
|
|
||||
// 参数:
|
|
||||
// - ctx: 上下文
|
|
||||
// - request: MCP 工具调用请求
|
|
||||
//
|
|
||||
// 返回值:
|
|
||||
// - result: MCP 工具调用结果
|
|
||||
// - err: 处理过程中的错误
|
|
||||
//
|
|
||||
// 异常:
|
|
||||
// - 无(不主动 panic)
|
|
||||
func (this *tool_spotify_music_playlist) Handl(ctx context.Context, request mcp.CallToolRequest) (result *mcp.CallToolResult, err error) { |
|
||||
var ( |
|
||||
offset, limit int |
|
||||
query string |
|
||||
data []spotify.SimplePlaylist |
|
||||
musics []spotify.PlaylistItem |
|
||||
card *SpotifyMusicListCard |
|
||||
context string |
|
||||
) |
|
||||
limit = request.GetInt("count", 10) |
|
||||
offset = request.GetInt("offset", 0) |
|
||||
if query, err = request.RequireString("query"); err != nil { |
|
||||
return |
|
||||
} |
|
||||
if data, err = lgspotify.SearchMusicList(ctx, query, 1, offset); err == nil { |
|
||||
musics, err = lgspotify.SearchMusicListDetails(ctx, data[0].ID, limit, 0) |
|
||||
} |
|
||||
|
|
||||
if err != nil { |
|
||||
context = "播放音乐失败" |
|
||||
log.Errorf("获取天气信息失败:%v", err) |
|
||||
return |
|
||||
} else { |
|
||||
card = &SpotifyMusicListCard{Musics: make([]SpotifyMusicCard, 0)} |
|
||||
for _, v := range musics { |
|
||||
item := SpotifyMusicCard{ |
|
||||
Name: v.Track.Track.Name, |
|
||||
Images: v.Track.Track.Album.Images[0].URL, |
|
||||
Url: string(v.Track.Track.URI), |
|
||||
Sgener: v.Track.Track.Album.Name, |
|
||||
} |
|
||||
card.Musics = append(card.Musics, item) |
|
||||
} |
|
||||
|
|
||||
context = utils.ToString(map[string]interface{}{ |
|
||||
"card_musiclist": card, |
|
||||
"result": fmt.Sprintf("正在为你播放音乐列表:%s", data[0].Name), |
|
||||
}) |
|
||||
log.Debug("播放音乐成功", |
|
||||
log.Field{Key: "query", Value: query}, |
|
||||
log.Field{Key: "card", Value: card}, |
|
||||
) |
|
||||
} |
|
||||
|
|
||||
result = &mcp.CallToolResult{ |
|
||||
Content: []mcp.Content{ |
|
||||
mcp.TextContent{ |
|
||||
Type: "text", |
|
||||
Text: context, |
|
||||
}, |
|
||||
}, |
|
||||
} |
|
||||
// result.Meta = map[string]interface{}{
|
|
||||
// "card_spotifylist": card,
|
|
||||
// }
|
|
||||
return |
|
||||
} |
|
||||
@ -1,79 +0,0 @@ |
|||||
package timer |
|
||||
|
|
||||
import ( |
|
||||
"context" |
|
||||
"yunyan/comm" |
|
||||
redissys "yunyan/lego/sys/redis" |
|
||||
"yunyan/modules" |
|
||||
"fmt" |
|
||||
"time" |
|
||||
|
|
||||
"yunyan/lego/base" |
|
||||
"yunyan/lego/core" |
|
||||
"yunyan/lego/sys/cron" |
|
||||
) |
|
||||
|
|
||||
type uselogTimer struct { |
|
||||
modules.MCompHttpGate |
|
||||
service base.IRPCXService |
|
||||
module *Timer |
|
||||
options *Options |
|
||||
} |
|
||||
|
|
||||
func (this *uselogTimer) Init(service core.IService, module core.IModule, comp core.IModuleComp, options core.IModuleOptions) (err error) { |
|
||||
this.MCompHttpGate.Init(service, module, comp, options) |
|
||||
this.service = service.(base.IRPCXService) |
|
||||
this.module = module.(*Timer) |
|
||||
this.options = options.(*Options) |
|
||||
|
|
||||
return |
|
||||
} |
|
||||
|
|
||||
func (this *uselogTimer) Start() (err error) { |
|
||||
err = this.MCompHttpGate.Start() |
|
||||
//凌晨1分1秒执行
|
|
||||
cron.AddFunc("1 1 0 * * ?", this.timer) |
|
||||
return |
|
||||
} |
|
||||
|
|
||||
func (this *uselogTimer) timer() { |
|
||||
fmt.Println("开始读取 log:* 数据...", time.Now().Format("15:04:05")) |
|
||||
ctx := context.Background() |
|
||||
// 使用 SCAN 而非 KEYS(更安全)
|
|
||||
var ( |
|
||||
cursor uint64 |
|
||||
datas []map[string]string = make([]map[string]string, 0) |
|
||||
) |
|
||||
for { |
|
||||
keys, newCursor, err := redissys.Conn().Scan(ctx, cursor, redissys.RKey(fmt.Sprintf("%s:*", comm.TableUseRecordLog)), 100).Result() |
|
||||
if err != nil { |
|
||||
fmt.Println("Scan error:", err) |
|
||||
break |
|
||||
} |
|
||||
|
|
||||
for _, key := range keys { |
|
||||
data, err := redissys.Conn().HGetAll(ctx, key).Result() |
|
||||
if err != nil { |
|
||||
fmt.Println("HGETALL error:", err) |
|
||||
continue |
|
||||
} |
|
||||
datas = append(datas, data) |
|
||||
// ✅ 这里处理数据,比如打印 / 存储到 DB
|
|
||||
fmt.Printf("处理 %s: %+v\n", key, data) |
|
||||
|
|
||||
// 删除日志 key
|
|
||||
if err := redissys.Conn().Del(ctx, key).Err(); err != nil { |
|
||||
fmt.Println("DEL error:", err) |
|
||||
} else { |
|
||||
fmt.Printf("已删除 %s\n", key) |
|
||||
} |
|
||||
} |
|
||||
|
|
||||
if newCursor == 0 { |
|
||||
break |
|
||||
} |
|
||||
cursor = newCursor |
|
||||
} |
|
||||
|
|
||||
fmt.Println("日志处理完成。") |
|
||||
} |
|
||||
@ -1,32 +0,0 @@ |
|||||
package sts |
|
||||
|
|
||||
type ( |
|
||||
Authorization struct { |
|
||||
Expiration string |
|
||||
AccessKeyId string |
|
||||
AccessKeySecret string |
|
||||
SecurityToken string |
|
||||
} |
|
||||
ISys interface { |
|
||||
//roleArn:角色ARN。
|
|
||||
AssumeRole(roleArn, roleSessionName string) (auth *Authorization, err error) |
|
||||
} |
|
||||
) |
|
||||
|
|
||||
var ( |
|
||||
defsys ISys |
|
||||
) |
|
||||
|
|
||||
func OnInit(config map[string]interface{}, option ...Option) (err error) { |
|
||||
defsys, err = newSys(newOptions(config, option...)) |
|
||||
return |
|
||||
} |
|
||||
|
|
||||
func NewSys(option ...Option) (sys ISys, err error) { |
|
||||
sys, err = newSys(newOptionsByOption(option...)) |
|
||||
return |
|
||||
} |
|
||||
|
|
||||
func AssumeRole(roleArn, roleSessionName string) (auth *Authorization, err error) { |
|
||||
return defsys.AssumeRole(roleArn, roleSessionName) |
|
||||
} |
|
||||
@ -1,52 +0,0 @@ |
|||||
package sts |
|
||||
|
|
||||
import ( |
|
||||
"yunyan/lego/utils/mapstructure" |
|
||||
) |
|
||||
|
|
||||
type Option func(*Options) |
|
||||
type Options struct { |
|
||||
RegionId string |
|
||||
AccessKeyId string |
|
||||
AccessKeySecret string |
|
||||
} |
|
||||
|
|
||||
// 注册点 cn-shenzhen
|
|
||||
func SetRegionId(v string) Option { |
|
||||
return func(o *Options) { |
|
||||
o.RegionId = v |
|
||||
} |
|
||||
} |
|
||||
|
|
||||
// 头像->访问控制->用户->用户 AccessKey
|
|
||||
func SetAccessKeyId(v string) Option { |
|
||||
return func(o *Options) { |
|
||||
o.AccessKeyId = v |
|
||||
} |
|
||||
} |
|
||||
|
|
||||
// 头像->访问控制->用户->用户 AccessKeySecret
|
|
||||
func SetAccessKeySecret(v string) Option { |
|
||||
return func(o *Options) { |
|
||||
o.AccessKeySecret = v |
|
||||
} |
|
||||
} |
|
||||
|
|
||||
func newOptions(config map[string]interface{}, opts ...Option) Options { |
|
||||
options := Options{} |
|
||||
if config != nil { |
|
||||
mapstructure.Decode(config, &options) |
|
||||
} |
|
||||
for _, o := range opts { |
|
||||
o(&options) |
|
||||
} |
|
||||
return options |
|
||||
} |
|
||||
|
|
||||
func newOptionsByOption(opts ...Option) Options { |
|
||||
options := Options{} |
|
||||
for _, o := range opts { |
|
||||
o(&options) |
|
||||
} |
|
||||
return options |
|
||||
} |
|
||||
@ -1,49 +0,0 @@ |
|||||
package sts |
|
||||
|
|
||||
import ( |
|
||||
"github.com/aliyun/alibaba-cloud-sdk-go/services/sts" |
|
||||
) |
|
||||
|
|
||||
func newSys(options Options) (sys *STS, err error) { |
|
||||
sys = &STS{options: options} |
|
||||
err = sys.init() |
|
||||
return |
|
||||
} |
|
||||
|
|
||||
type STS struct { |
|
||||
options Options |
|
||||
client *sts.Client |
|
||||
} |
|
||||
|
|
||||
func (this *STS) init() (err error) { |
|
||||
if this.client, err = sts.NewClientWithAccessKey(this.options.RegionId, this.options.AccessKeyId, this.options.AccessKeySecret); err != nil { |
|
||||
return |
|
||||
} |
|
||||
|
|
||||
return |
|
||||
} |
|
||||
|
|
||||
func (this *STS) AssumeRole(roleArn, roleSessionName string) (auth *Authorization, err error) { |
|
||||
var ( |
|
||||
request *sts.AssumeRoleRequest |
|
||||
response *sts.AssumeRoleResponse |
|
||||
) |
|
||||
//构建请求对象。
|
|
||||
request = sts.CreateAssumeRoleRequest() |
|
||||
request.Scheme = "https" |
|
||||
//设置参数。关于参数含义和设置方法,请参见《API参考》。
|
|
||||
request.RoleArn = roleArn |
|
||||
request.RoleSessionName = roleSessionName |
|
||||
//发起请求,并得到响应。
|
|
||||
response, err = this.client.AssumeRole(request) |
|
||||
if err != nil { |
|
||||
return |
|
||||
} |
|
||||
auth = &Authorization{ |
|
||||
Expiration: response.Credentials.Expiration, |
|
||||
AccessKeyId: response.Credentials.AccessKeyId, |
|
||||
AccessKeySecret: response.Credentials.AccessKeySecret, |
|
||||
SecurityToken: response.Credentials.SecurityToken, |
|
||||
} |
|
||||
return |
|
||||
} |
|
||||
@ -1,30 +0,0 @@ |
|||||
package sts_test |
|
||||
|
|
||||
import ( |
|
||||
"fmt" |
|
||||
"testing" |
|
||||
|
|
||||
"yunyan/sys/aliyun/sts" |
|
||||
) |
|
||||
|
|
||||
// oss: #对象存储
|
|
||||
// Endpoint: https://oss-ap-southeast-1.aliyuncs.com
|
|
||||
// AccessKeyId: xxxxxxxxxxxxx
|
|
||||
// AccessKeySecret: xxxxxxxxxxxxxxxxxxx
|
|
||||
// BucketName: dpmobj
|
|
||||
|
|
||||
func Test_STS(t *testing.T) { |
|
||||
sys, err := sts.NewSys( |
|
||||
sts.SetRegionId("cn-shenzhen"), |
|
||||
sts.SetAccessKeyId("xxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxx"), |
|
||||
sts.SetAccessKeySecret("xxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxx"), |
|
||||
) |
|
||||
if err != nil { |
|
||||
fmt.Printf("初始化OSS 系统失败 err:%v", err) |
|
||||
return |
|
||||
} else { |
|
||||
fmt.Printf("初始化OSS 系统成功") |
|
||||
auth, err := sys.AssumeRole("xxxxxxxxxxxxxxxxxxxxxxxxxxx", "SessionTest") |
|
||||
fmt.Printf("初始化OSS AssumeRole auth:%+v err:%v", auth, err) |
|
||||
} |
|
||||
} |
|
||||
@ -1,84 +0,0 @@ |
|||||
package axml |
|
||||
|
|
||||
import ( |
|
||||
"regexp" |
|
||||
"strings" |
|
||||
) |
|
||||
|
|
||||
func newSys(options Options) (sys ISys, err error) { |
|
||||
sys = &AXml{} |
|
||||
return |
|
||||
} |
|
||||
|
|
||||
// XMLItem 表示一个通用的 XML 数据项
|
|
||||
type XMLItem struct { |
|
||||
TagName string // 标签名
|
|
||||
Attributes map[string]string // 属性键值对
|
|
||||
Content string // 标签内容
|
|
||||
} |
|
||||
type AXml struct { |
|
||||
source string |
|
||||
buffer string |
|
||||
items []*XMLItem |
|
||||
} |
|
||||
|
|
||||
func (this *AXml) Add(data string) (err error) { |
|
||||
this.source += data |
|
||||
this.buffer += data |
|
||||
this.processData() |
|
||||
return |
|
||||
} |
|
||||
|
|
||||
func (this *AXml) Get(name string) (items []*XMLItem) { |
|
||||
items = make([]*XMLItem, 0) |
|
||||
for _, v := range this.items { |
|
||||
if v.TagName == name { |
|
||||
items = append(items, v) |
|
||||
} |
|
||||
} |
|
||||
return |
|
||||
} |
|
||||
|
|
||||
// 解析数据
|
|
||||
func (this *AXml) processData() { |
|
||||
// 匹配完整的 <tag attr="value">content</tag> 模式
|
|
||||
// 捕获标签名、属性和内容
|
|
||||
re := regexp.MustCompile(`<([a-zA-Z][^>]*)>((?:[^<]|<[^/])*?)</([a-zA-Z][^>]*)>`) |
|
||||
matches := re.FindAllStringSubmatch(this.buffer, -1) |
|
||||
|
|
||||
// var items []XMLItem
|
|
||||
for _, match := range matches { |
|
||||
if len(match) != 4 { |
|
||||
continue |
|
||||
} |
|
||||
// match[0] 是完整匹配,match[1] 是标签和属性部分,match[2] 是内容
|
|
||||
tagAndAttrs := match[1] |
|
||||
content := match[2] |
|
||||
|
|
||||
// 解析标签名和属性
|
|
||||
attrRe := regexp.MustCompile(`([a-zA-Z]+)(?: *= *"([^"]*)")?`) |
|
||||
attrMatches := attrRe.FindAllStringSubmatch(tagAndAttrs, -1) |
|
||||
|
|
||||
item := &XMLItem{ |
|
||||
Attributes: make(map[string]string), |
|
||||
} |
|
||||
for _, attrMatch := range attrMatches { |
|
||||
if len(attrMatch) >= 2 { |
|
||||
if item.TagName == "" { |
|
||||
item.TagName = attrMatch[1] // 第一个捕获组是标签名
|
|
||||
} |
|
||||
if len(attrMatch) == 3 && attrMatch[2] != "" { |
|
||||
item.Attributes[attrMatch[1]] = attrMatch[2] // 属性名和值
|
|
||||
} |
|
||||
} |
|
||||
} |
|
||||
item.Content = strings.TrimSpace(content) |
|
||||
|
|
||||
if item.TagName != "" { |
|
||||
this.items = append(this.items, item) |
|
||||
// 从缓冲区移除已解析的部分
|
|
||||
this.buffer = this.buffer[strings.Index(this.buffer, match[0])+len(match[0]):] |
|
||||
} |
|
||||
} |
|
||||
|
|
||||
} |
|
||||
@ -1,30 +0,0 @@ |
|||||
package axml |
|
||||
|
|
||||
type ( |
|
||||
ISys interface { |
|
||||
Add(data string) (err error) |
|
||||
Get(name string) (items []*XMLItem) |
|
||||
} |
|
||||
) |
|
||||
|
|
||||
var ( |
|
||||
defsys ISys |
|
||||
) |
|
||||
|
|
||||
func OnInit(config map[string]interface{}, option ...Option) (err error) { |
|
||||
defsys, err = newSys(newOptions(config, option...)) |
|
||||
return |
|
||||
} |
|
||||
|
|
||||
func NewSys(option ...Option) (sys ISys, err error) { |
|
||||
sys, err = newSys(newOptionsByOption(option...)) |
|
||||
return |
|
||||
} |
|
||||
|
|
||||
func Add(data string) error { |
|
||||
return defsys.Add(data) |
|
||||
} |
|
||||
|
|
||||
func Get(name string) (items []*XMLItem) { |
|
||||
return defsys.Get(name) |
|
||||
} |
|
||||
@ -1,37 +0,0 @@ |
|||||
package axml |
|
||||
|
|
||||
import ( |
|
||||
"yunyan/lego/sys/log" |
|
||||
"yunyan/lego/utils/mapstructure" |
|
||||
) |
|
||||
|
|
||||
type Option func(*Options) |
|
||||
type Options struct { |
|
||||
Debug bool //日志是否开启
|
|
||||
Log log.ILogger |
|
||||
} |
|
||||
|
|
||||
func newOptions(config map[string]interface{}, opts ...Option) Options { |
|
||||
options := Options{} |
|
||||
if config != nil { |
|
||||
mapstructure.Decode(config, &options) |
|
||||
} |
|
||||
for _, o := range opts { |
|
||||
o(&options) |
|
||||
} |
|
||||
if options.Log == nil { |
|
||||
options.Log = log.NewTurnlog(options.Debug, log.Clone("sys.deepseek", 3)) |
|
||||
} |
|
||||
return options |
|
||||
} |
|
||||
|
|
||||
func newOptionsByOption(opts ...Option) Options { |
|
||||
options := Options{} |
|
||||
for _, o := range opts { |
|
||||
o(&options) |
|
||||
} |
|
||||
if options.Log == nil { |
|
||||
options.Log = log.NewTurnlog(options.Debug, log.Clone("sys.deepseek", 3)) |
|
||||
} |
|
||||
return options |
|
||||
} |
|
||||
@ -1,56 +0,0 @@ |
|||||
package coze |
|
||||
|
|
||||
import "context" |
|
||||
|
|
||||
type ( |
|
||||
Message struct { |
|
||||
Role string `json:"role"` |
|
||||
Content string `json:"content"` |
|
||||
} |
|
||||
//回应流式切片
|
|
||||
ChatResponseChoice struct { |
|
||||
Role string `json:"role"` |
|
||||
Content string `json:"content,omitempty"` |
|
||||
Meta map[string]interface{} `json:"meta,omitempty"` |
|
||||
} |
|
||||
// ToolResponseCard 通用卡片结构(不解析底层元素细节)
|
|
||||
CardContent struct { |
|
||||
CardType int `json:"card_type"` |
|
||||
TemplateURL string `json:"template_url"` |
|
||||
TemplateID int64 `json:"template_id"` |
|
||||
ResponseForModel string `json:"response_for_model"` |
|
||||
ContentType int `json:"content_type"` |
|
||||
Data string `json:"data"` // 卡片结构JSON字符串
|
|
||||
Variables map[string]any `json:"variables"` // 变量数据
|
|
||||
InfoInCard string `json:"info_in_card"` // 备用数据
|
|
||||
ResponseType string `json:"response_type"` |
|
||||
XProperties map[string]any `json:"x_properties"` |
|
||||
} |
|
||||
CardContentData struct { |
|
||||
Variables map[string]CardData `json:"variables"` //卡片数据
|
|
||||
} |
|
||||
CardData struct { |
|
||||
ID string `json:"ID"` |
|
||||
Name string `json:"name"` |
|
||||
DefaultValue interface{} `json:"defaultValue"` |
|
||||
} |
|
||||
ISys interface { |
|
||||
ChatForSteams(ctx context.Context, uid string, customVariables map[string]string, messages []Message, choiceChan chan *ChatResponseChoice) (err error) |
|
||||
} |
|
||||
) |
|
||||
|
|
||||
var defsys ISys |
|
||||
|
|
||||
func OnInit(config map[string]interface{}, option ...Option) (err error) { |
|
||||
defsys, err = newSys(newOptions(config, option...)) |
|
||||
return |
|
||||
} |
|
||||
|
|
||||
func NewSys(option ...Option) (sys ISys, err error) { |
|
||||
sys, err = newSys(newOptionsByOption(option...)) |
|
||||
return |
|
||||
} |
|
||||
|
|
||||
func ChatForSteams(ctx context.Context, uid string, customVariables map[string]string, messages []Message, choiceChan chan *ChatResponseChoice) (err error) { |
|
||||
return defsys.ChatForSteams(ctx, uid, customVariables, messages, choiceChan) |
|
||||
} |
|
||||
@ -1,132 +0,0 @@ |
|||||
package coze |
|
||||
|
|
||||
import ( |
|
||||
"context" |
|
||||
"encoding/json" |
|
||||
"errors" |
|
||||
"fmt" |
|
||||
"io" |
|
||||
"net/http" |
|
||||
"time" |
|
||||
|
|
||||
"github.com/coze-dev/coze-go" |
|
||||
) |
|
||||
|
|
||||
func newSys(options Options) (sys *Coze, err error) { |
|
||||
sys = &Coze{ |
|
||||
options: options, |
|
||||
} |
|
||||
// sys.Init()
|
|
||||
authCli := coze.NewTokenAuth(options.Token) |
|
||||
sys.api = coze.NewCozeAPI(authCli, coze.WithBaseURL(coze.CnBaseURL)) |
|
||||
return |
|
||||
} |
|
||||
|
|
||||
type Coze struct { |
|
||||
options Options |
|
||||
api coze.CozeAPI |
|
||||
} |
|
||||
|
|
||||
func (this *Coze) Init() { |
|
||||
// Get an access_token through personal access token or oauth.
|
|
||||
// token := os.Getenv("COZE_API_TOKEN")
|
|
||||
authCli := coze.NewTokenAuth(this.options.Token) |
|
||||
|
|
||||
// 1. Initialize with default configuration
|
|
||||
cozeCli1 := coze.NewCozeAPI(authCli) |
|
||||
fmt.Println("client 1:", cozeCli1) |
|
||||
|
|
||||
// 2. Initialize with custom base URL
|
|
||||
// cozeAPIBase := os.Getenv("COZE_API_BASE")
|
|
||||
cozeCli2 := coze.NewCozeAPI(authCli, coze.WithBaseURL(coze.ComBaseURL)) |
|
||||
fmt.Println("client 2:", cozeCli2) |
|
||||
|
|
||||
// 3. Initialize with custom HTTP client
|
|
||||
customClient := &http.Client{ |
|
||||
Timeout: 30 * time.Second, |
|
||||
Transport: &http.Transport{ |
|
||||
MaxIdleConns: 100, |
|
||||
MaxIdleConnsPerHost: 100, |
|
||||
IdleConnTimeout: 90 * time.Second, |
|
||||
}, |
|
||||
} |
|
||||
cozeCli3 := coze.NewCozeAPI(authCli, |
|
||||
coze.WithBaseURL(coze.ComBaseURL), |
|
||||
coze.WithHttpClient(customClient), |
|
||||
) |
|
||||
fmt.Println("client 3:", cozeCli3) |
|
||||
} |
|
||||
|
|
||||
func (this *Coze) ChatForSteams(ctx context.Context, uid string, customVariables map[string]string, messages []Message, choiceChan chan *ChatResponseChoice) (err error) { |
|
||||
var ( |
|
||||
msgs []*coze.Message = make([]*coze.Message, len(messages)) |
|
||||
resp coze.Stream[coze.ChatEvent] |
|
||||
) |
|
||||
|
|
||||
for i, v := range messages { |
|
||||
msgs[i] = &coze.Message{ |
|
||||
Role: coze.MessageRoleUser, |
|
||||
Content: v.Content, |
|
||||
} |
|
||||
} |
|
||||
req := &coze.CreateChatsReq{ |
|
||||
BotID: this.options.BotID, |
|
||||
UserID: uid, |
|
||||
Messages: msgs, |
|
||||
CustomVariables: customVariables, |
|
||||
} |
|
||||
ctx, cancel := context.WithTimeout(context.Background(), 10*time.Second) |
|
||||
defer cancel() |
|
||||
defer close(choiceChan) |
|
||||
if resp, err = this.api.Chat.Stream(ctx, req); err != nil { |
|
||||
this.options.Log.Errorln(err) |
|
||||
return |
|
||||
} |
|
||||
defer resp.Close() |
|
||||
for { |
|
||||
event, err := resp.Recv() |
|
||||
if errors.Is(err, io.EOF) { |
|
||||
break |
|
||||
} |
|
||||
if err != nil { |
|
||||
this.options.Log.Errorln(err) |
|
||||
break |
|
||||
} |
|
||||
// fmt.Printf("结果:Event: %s Message:%+v \n", event.Event, event.Message)
|
|
||||
if event.Event == coze.ChatEventConversationMessageDelta { |
|
||||
choiceChan <- &ChatResponseChoice{ |
|
||||
Role: "ai", |
|
||||
Content: event.Message.Content, |
|
||||
} |
|
||||
} else if event.Event == coze.ChatEventConversationMessageCompleted { |
|
||||
if event.Message.Type == coze.MessageTypeToolResponse { |
|
||||
// fmt.Printf("结果:Event: %s Message:%+v \n\n\n", event.Event, event.Message)
|
|
||||
card := CardContent{} |
|
||||
if err = json.Unmarshal([]byte(event.Message.Content), &card); err != nil { |
|
||||
this.options.Log.Errorln(err) |
|
||||
} else { |
|
||||
data := CardContentData{} |
|
||||
if card.Data != "" { |
|
||||
if err = json.Unmarshal([]byte(card.Data), &data); err != nil { |
|
||||
fmt.Println(card.Data) |
|
||||
this.options.Log.Errorln(card.Data, err) |
|
||||
} else { |
|
||||
if len(data.Variables) > 0 { |
|
||||
for _, v := range data.Variables { |
|
||||
choiceChan <- &ChatResponseChoice{ |
|
||||
Role: "card", |
|
||||
Meta: map[string]interface{}{ |
|
||||
v.Name: v.DefaultValue, |
|
||||
}, |
|
||||
} |
|
||||
} |
|
||||
|
|
||||
} |
|
||||
} |
|
||||
} |
|
||||
} |
|
||||
} |
|
||||
} |
|
||||
} |
|
||||
return |
|
||||
} |
|
||||
@ -1,49 +0,0 @@ |
|||||
package coze |
|
||||
|
|
||||
import ( |
|
||||
"yunyan/lego/sys/log" |
|
||||
"yunyan/lego/utils/mapstructure" |
|
||||
) |
|
||||
|
|
||||
type Option func(*Options) |
|
||||
type Options struct { |
|
||||
Debug bool //日志是否开启
|
|
||||
Log log.ILogger |
|
||||
Token string |
|
||||
BotID string |
|
||||
} |
|
||||
|
|
||||
func SetToken(v string) Option { |
|
||||
return func(o *Options) { |
|
||||
o.Token = v |
|
||||
} |
|
||||
} |
|
||||
func SetBotID(v string) Option { |
|
||||
return func(o *Options) { |
|
||||
o.BotID = v |
|
||||
} |
|
||||
} |
|
||||
func newOptions(config map[string]interface{}, opts ...Option) Options { |
|
||||
options := Options{} |
|
||||
if config != nil { |
|
||||
mapstructure.Decode(config, &options) |
|
||||
} |
|
||||
for _, o := range opts { |
|
||||
o(&options) |
|
||||
} |
|
||||
if options.Log == nil { |
|
||||
options.Log = log.NewTurnlog(options.Debug, log.Clone("sys.coze", 3)) |
|
||||
} |
|
||||
return options |
|
||||
} |
|
||||
|
|
||||
func newOptionsByOption(opts ...Option) Options { |
|
||||
options := Options{} |
|
||||
for _, o := range opts { |
|
||||
o(&options) |
|
||||
} |
|
||||
if options.Log == nil { |
|
||||
options.Log = log.NewTurnlog(options.Debug, log.Clone("sys.coze", 3)) |
|
||||
} |
|
||||
return options |
|
||||
} |
|
||||
@ -1,38 +0,0 @@ |
|||||
package coze_test |
|
||||
|
|
||||
import ( |
|
||||
|
|
||||
//"lego_bighealth/sys/coze"
|
|
||||
|
|
||||
"context" |
|
||||
"yunyan/sys/coze" |
|
||||
"fmt" |
|
||||
"testing" |
|
||||
) |
|
||||
|
|
||||
func Test_Sys_Chat(t *testing.T) { |
|
||||
if sys, err := coze.NewSys( |
|
||||
coze.SetToken("pat_wOdiIjLLWuqHxIla03qr81FqWmestCGe1WIfTGUALQhuPbtm3vmntQbtfds9SikP"), |
|
||||
coze.SetBotID("7486826364478898226"), |
|
||||
); err != nil { |
|
||||
fmt.Printf("Sys Init err:%v", err) |
|
||||
} else { |
|
||||
result := make(chan *coze.ChatResponseChoice, 10) |
|
||||
go func() { |
|
||||
err := sys.ChatForSteams(context.Background(), "liwei1dao", map[string]string{"sys_lon_lat": "116.321669,39.985266"}, []coze.Message{ |
|
||||
{ |
|
||||
Role: "user", |
|
||||
Content: "给我导航到大新地铁站", |
|
||||
}, |
|
||||
}, result) |
|
||||
if err != nil { |
|
||||
fmt.Printf(" err:%v", err) |
|
||||
return |
|
||||
} |
|
||||
}() |
|
||||
|
|
||||
for v := range result { |
|
||||
fmt.Printf("result:%v\n", v) |
|
||||
} |
|
||||
} |
|
||||
} |
|
||||
@ -1,119 +0,0 @@ |
|||||
package deepseek |
|
||||
|
|
||||
type ( |
|
||||
ISys interface { |
|
||||
Chat(msgs []Message) (result []string, err error) |
|
||||
} |
|
||||
//消息
|
|
||||
Message struct { |
|
||||
//发起消息的角色
|
|
||||
//Possible values: [system,user,assistant,tool];
|
|
||||
Role string `json:"role"` |
|
||||
//消息的内容
|
|
||||
Content string `json:"content"` |
|
||||
//是否访问网络
|
|
||||
Web bool `json:"web"` |
|
||||
//可以选填的参与者的名称,为模型提供信息以区分相同角色的参与者。
|
|
||||
Name string `json:"name"` |
|
||||
//(Beta) 设置此参数为 true,来强制模型在其回答中以此 assistant 消息中提供的前缀内容开始。
|
|
||||
//您必须设置 base_url="https://api.deepseek.com/beta" 来使用此功能
|
|
||||
Prefix bool `json:"prefix"` |
|
||||
//(Beta) 用于 deepseek-reasoner 模型在对话前缀续写功能下,作为最后一条 assistant 思维链内容的输入。使用此功能时,prefix 参数必须设置为 true。
|
|
||||
ReasoningContent string `json:"reasoning_content"` |
|
||||
//此消息所响应的 tool call 的 ID。
|
|
||||
ToolCallId string `json:"tool_call_id"` |
|
||||
} |
|
||||
//回应控制
|
|
||||
ResponesFormat struct { |
|
||||
Type string `json:"type"` |
|
||||
} |
|
||||
//流式参数
|
|
||||
StreaOptions struct { |
|
||||
//如果设置为 true,在流式消息最后的 data: [DONE] 之前将会传输一个额外的块。此块上的 usage 字段显示整个请求的 token 使用统计信息,而 choices 字段将始终是一个空数组。所有其他块也将包含一个 usage 字段,但其值为 null。
|
|
||||
IncludeUsage bool `json:"include_usage"` |
|
||||
} |
|
||||
//工具
|
|
||||
Tool struct { |
|
||||
Type string `json:"type"` //工具名称
|
|
||||
Function Function `json:"function"` //工具名称
|
|
||||
} |
|
||||
//工具
|
|
||||
Function struct { |
|
||||
Name string `json:"name"` //方法名称
|
|
||||
Description string `json:"description"` //方法描述
|
|
||||
Parameters Parameters `json:"parameters"` //参数列表
|
|
||||
} |
|
||||
//方法调用 参数列表
|
|
||||
Parameters struct { |
|
||||
Type string `json:"type"` //传参勒烯
|
|
||||
Properties map[string]Propertie `json:"properties"` //参数列表
|
|
||||
Required []string `json:"required"` //必传参数列表
|
|
||||
} |
|
||||
//参数
|
|
||||
Propertie struct { |
|
||||
Type string `json:"type"` //参数类型
|
|
||||
Description string `json:"description"` //参数描述
|
|
||||
} |
|
||||
ChatRequest struct { |
|
||||
//对话的消息列表。
|
|
||||
Messages []Message `json:"messages"` |
|
||||
//Possible values: [deepseek-chat, deepseek-reasoner];
|
|
||||
//使用的模型的 ID。您可以使用 deepseek-chat。
|
|
||||
Model string `json:"model"` |
|
||||
//Possible values: >= -2 and <= 2;
|
|
||||
//Default value: 0;
|
|
||||
//介于 -2.0 和 2.0 之间的数字。如果该值为正,那么新 token 会根据其在已有文本中的出现频率受到相应的惩罚,降低模型重复相同内容的可能性。
|
|
||||
FrequencyPenalty string `json:"frequency_penalty"` |
|
||||
//Possible values: > 1;
|
|
||||
//介于 1 到 8192 间的整数,限制一次请求中模型生成 completion 的最大 token 数。输入 token 和输出 token 的总长度受模型的上下文长度的限制。
|
|
||||
//如未指定 max_tokens参数,默认使用 4096。
|
|
||||
MaxTokens int `json:"max_tokens"` |
|
||||
//Possible values: >= -2 and <= 2
|
|
||||
//Default value: 0
|
|
||||
//介于 -2.0 和 2.0 之间的数字。如果该值为正,那么新 token 会根据其是否已在已有文本中出现受到相应的惩罚,从而增加模型谈论新主题的可能性。
|
|
||||
PresencePenalty int `json:"presence_penalty"` |
|
||||
//一个 object,指定模型必须输出的格式。设置为 { "type": "json_object" } 以启用 JSON 模式,该模式保证模型生成的消息是有效的 JSON。
|
|
||||
//注意: 使用 JSON 模式时,你还必须通过系统或用户消息指示模型生成 JSON。否则,模型可能会生成不断的空白字符,直到生成达到令牌限制,从而导致请求长时间运行并显得“卡住”。此外,如果 finish_reason="length",这表示生成超过了 max_tokens 或对话超过了最大上下文长度,消息内容可能会被部分截断
|
|
||||
ResponesFormat ResponesFormat `json:"response_format"` |
|
||||
//一个 string 或最多包含 16 个 string 的 list,在遇到这些词时,API 将停止生成更多的 token。
|
|
||||
Stop []string `json:"stop"` |
|
||||
//如果设置为 True,将会以 SSE(server-sent events)的形式以流式发送消息增量。消息流以
|
|
||||
Stream bool `json:"stream"` |
|
||||
//流式输出相关选项。只有在 stream 参数为 true 时,才可设置此参数
|
|
||||
StreaOptions *StreaOptions `json:"stream_options"` |
|
||||
//Possible values: <= 2
|
|
||||
//Default value: 1
|
|
||||
//采样温度,介于 0 和 2 之间。更高的值,如 0.8,会使输出更随机,而更低的值,如 0.2,会使其更加集中和确定。 我们通常建议可以更改这个值或者更改 top_p,但不建议同时对两者进行修改。
|
|
||||
Temperature int `json:"temperature"` |
|
||||
//Possible values: <= 1
|
|
||||
//Default value: 1
|
|
||||
//作为调节采样温度的替代方案,模型会考虑前 top_p 概率的 token 的结果。所以 0.1 就意味着只有包括在最高 10% 概率中的 token 会被考虑。 我们通常建议修改这个值或者更改 temperature,但不建议同时对两者进行修改
|
|
||||
TopP int `json:"top_p"` |
|
||||
//模型可能会调用的 tool 的列表。目前,仅支持 function 作为工具。使用此参数来提供以 JSON 作为输入参数的 function 列表。最多支持 128 个 function。
|
|
||||
Tools []Tool |
|
||||
} |
|
||||
|
|
||||
ChatResponse struct { |
|
||||
// 根据实际的API响应定义结构体字段
|
|
||||
// 例如:
|
|
||||
Choices []struct { |
|
||||
Message Message `json:"message"` |
|
||||
} `json:"choices"` |
|
||||
} |
|
||||
) |
|
||||
|
|
||||
var defsys ISys |
|
||||
|
|
||||
func OnInit(config map[string]interface{}, option ...Option) (err error) { |
|
||||
defsys, err = newSys(newOptions(config, option...)) |
|
||||
return |
|
||||
} |
|
||||
|
|
||||
func NewSys(option ...Option) (sys ISys, err error) { |
|
||||
sys, err = newSys(newOptionsByOption(option...)) |
|
||||
return |
|
||||
} |
|
||||
|
|
||||
func Chat(msgs []Message) (result []string, err error) { |
|
||||
return defsys.Chat(msgs) |
|
||||
} |
|
||||
@ -1,86 +0,0 @@ |
|||||
package deepseek |
|
||||
|
|
||||
import ( |
|
||||
"bytes" |
|
||||
"encoding/json" |
|
||||
"fmt" |
|
||||
"io" |
|
||||
"net/http" |
|
||||
"sync" |
|
||||
) |
|
||||
|
|
||||
func newSys(options Options) (sys *Deepseek, err error) { |
|
||||
sys = &Deepseek{ |
|
||||
options: options, |
|
||||
pool: sync.Pool{ |
|
||||
New: func() interface{} { |
|
||||
return &http.Client{} |
|
||||
}, |
|
||||
}} |
|
||||
|
|
||||
return |
|
||||
} |
|
||||
|
|
||||
type Deepseek struct { |
|
||||
options Options |
|
||||
pool sync.Pool |
|
||||
} |
|
||||
|
|
||||
func (this *Deepseek) Chat(msgs []Message) (result []string, err error) { |
|
||||
var ( |
|
||||
jsonData []byte |
|
||||
body []byte |
|
||||
resp *http.Response |
|
||||
) |
|
||||
|
|
||||
url := "https://api.deepseek.com/v1/chat/completions" |
|
||||
chatRequest := ChatRequest{ |
|
||||
Model: "deepseek-chat", |
|
||||
Messages: msgs, |
|
||||
Stream: false, |
|
||||
} |
|
||||
|
|
||||
jsonData, err = json.Marshal(chatRequest) |
|
||||
if err != nil { |
|
||||
fmt.Println("Error marshalling JSON:", err) |
|
||||
return |
|
||||
} |
|
||||
|
|
||||
req, err := http.NewRequest("POST", url, bytes.NewBuffer(jsonData)) |
|
||||
if err != nil { |
|
||||
fmt.Println("Error creating request:", err) |
|
||||
return |
|
||||
} |
|
||||
|
|
||||
req.Header.Set("Content-Type", "application/json") |
|
||||
req.Header.Set("Authorization", "Bearer "+this.options.Appkey) |
|
||||
|
|
||||
client := this.pool.Get().(*http.Client) |
|
||||
defer this.pool.Put(client) |
|
||||
resp, err = client.Do(req) |
|
||||
if err != nil { |
|
||||
this.options.Log.Errorln("Error making request:", err) |
|
||||
return |
|
||||
} |
|
||||
defer resp.Body.Close() |
|
||||
|
|
||||
body, err = io.ReadAll(resp.Body) |
|
||||
if err != nil { |
|
||||
fmt.Println("Error reading response body:", err) |
|
||||
return |
|
||||
} |
|
||||
|
|
||||
var chatResponse ChatResponse |
|
||||
err = json.Unmarshal(body, &chatResponse) |
|
||||
if err != nil { |
|
||||
fmt.Println("Error unmarshalling response:", err) |
|
||||
return |
|
||||
} |
|
||||
result = make([]string, len(chatResponse.Choices)) |
|
||||
// 输出响应结果
|
|
||||
for _, choice := range chatResponse.Choices { |
|
||||
// fmt.Println("Response:", choice.Message.Content)
|
|
||||
result = append(result, choice.Message.Content) |
|
||||
} |
|
||||
return |
|
||||
} |
|
||||
@ -1,44 +0,0 @@ |
|||||
package deepseek |
|
||||
|
|
||||
import ( |
|
||||
"yunyan/lego/sys/log" |
|
||||
"yunyan/lego/utils/mapstructure" |
|
||||
) |
|
||||
|
|
||||
type Option func(*Options) |
|
||||
type Options struct { |
|
||||
Debug bool //日志是否开启
|
|
||||
Log log.ILogger |
|
||||
Appkey string |
|
||||
} |
|
||||
|
|
||||
func SetAppkey(v string) Option { |
|
||||
return func(o *Options) { |
|
||||
o.Appkey = v |
|
||||
} |
|
||||
} |
|
||||
|
|
||||
func newOptions(config map[string]interface{}, opts ...Option) Options { |
|
||||
options := Options{} |
|
||||
if config != nil { |
|
||||
mapstructure.Decode(config, &options) |
|
||||
} |
|
||||
for _, o := range opts { |
|
||||
o(&options) |
|
||||
} |
|
||||
if options.Log == nil { |
|
||||
options.Log = log.NewTurnlog(options.Debug, log.Clone("sys.deepseek", 3)) |
|
||||
} |
|
||||
return options |
|
||||
} |
|
||||
|
|
||||
func newOptionsByOption(opts ...Option) Options { |
|
||||
options := Options{} |
|
||||
for _, o := range opts { |
|
||||
o(&options) |
|
||||
} |
|
||||
if options.Log == nil { |
|
||||
options.Log = log.NewTurnlog(options.Debug, log.Clone("sys.deepseek", 3)) |
|
||||
} |
|
||||
return options |
|
||||
} |
|
||||
@ -1,157 +0,0 @@ |
|||||
package deepseek_test |
|
||||
|
|
||||
import ( |
|
||||
"bytes" |
|
||||
"context" |
|
||||
"yunyan/sys/deepseek" |
|
||||
"encoding/json" |
|
||||
"fmt" |
|
||||
"net/http" |
|
||||
"testing" |
|
||||
"time" |
|
||||
) |
|
||||
|
|
||||
const ( |
|
||||
deepseekAPIURL = "https://api.deepseek.com/v1/chat/completions" |
|
||||
) |
|
||||
|
|
||||
type DeepSeekClient struct { |
|
||||
apiKey string |
|
||||
httpClient *http.Client |
|
||||
} |
|
||||
|
|
||||
func NewDeepSeekClient(apiKey string) *DeepSeekClient { |
|
||||
return &DeepSeekClient{ |
|
||||
apiKey: apiKey, |
|
||||
httpClient: &http.Client{Timeout: 30 * time.Second}, |
|
||||
} |
|
||||
} |
|
||||
|
|
||||
// 增强版请求结构(假设支持搜索参数)
|
|
||||
type ChatRequest struct { |
|
||||
Model string `json:"model"` // 指定支持搜索的模型
|
|
||||
Messages []Message `json:"messages"` // 对话历史
|
|
||||
SearchConfig *SearchConfig `json:"search_config"` // 搜索配置
|
|
||||
Temperature float64 `json:"temperature,omitempty"` |
|
||||
} |
|
||||
|
|
||||
type SearchConfig struct { |
|
||||
Enable bool `json:"enable"` // 启用搜索
|
|
||||
SearchDepth int `json:"search_depth"` // 搜索深度
|
|
||||
RealTimeSearch bool `json:"real_time_search"` // 实时搜索
|
|
||||
ResultMaxTokens int `json:"result_max_tokens"` // 结果最大长度
|
|
||||
} |
|
||||
|
|
||||
type Message struct { |
|
||||
Role string `json:"role"` |
|
||||
Content string `json:"content"` |
|
||||
} |
|
||||
|
|
||||
type ChatResponse struct { |
|
||||
Choices []struct { |
|
||||
Message Message `json:"message"` |
|
||||
SearchResult []struct { // 假设返回包含搜索结果
|
|
||||
Title string `json:"title"` |
|
||||
Snippet string `json:"snippet"` |
|
||||
URL string `json:"url"` |
|
||||
} `json:"search_results,omitempty"` |
|
||||
} `json:"choices"` |
|
||||
} |
|
||||
|
|
||||
func (c *DeepSeekClient) ChatWithSearch(ctx context.Context, req ChatRequest) (*ChatResponse, error) { |
|
||||
reqBody, _ := json.Marshal(req) |
|
||||
|
|
||||
httpReq, _ := http.NewRequestWithContext(ctx, "POST", deepseekAPIURL, bytes.NewReader(reqBody)) |
|
||||
httpReq.Header.Set("Authorization", "Bearer "+c.apiKey) |
|
||||
httpReq.Header.Set("Content-Type", "application/json") |
|
||||
httpReq.Header.Set("Accept", "application/json") |
|
||||
|
|
||||
resp, err := c.httpClient.Do(httpReq) |
|
||||
if err != nil { |
|
||||
return nil, fmt.Errorf("API请求失败: %w", err) |
|
||||
} |
|
||||
defer resp.Body.Close() |
|
||||
|
|
||||
if resp.StatusCode != http.StatusOK { |
|
||||
return nil, fmt.Errorf("API返回异常状态码: %d", resp.StatusCode) |
|
||||
} |
|
||||
|
|
||||
var response ChatResponse |
|
||||
if err := json.NewDecoder(resp.Body).Decode(&response); err != nil { |
|
||||
return nil, fmt.Errorf("响应解析失败: %w", err) |
|
||||
} |
|
||||
|
|
||||
return &response, nil |
|
||||
} |
|
||||
|
|
||||
func Test_DeepSeek(t *testing.T) { |
|
||||
apiKey := "sk-3adfd188a3134e718bbf704f525aff17" |
|
||||
|
|
||||
client := NewDeepSeekClient(apiKey) |
|
||||
|
|
||||
// 构造包含搜索功能的请求
|
|
||||
request := ChatRequest{ |
|
||||
Model: "deepseek-chat", // 假设支持搜索的模型名称
|
|
||||
Messages: []Message{ |
|
||||
{ |
|
||||
Role: "user", |
|
||||
Content: "深圳今天的天气如何?", |
|
||||
}, |
|
||||
}, |
|
||||
SearchConfig: &SearchConfig{ |
|
||||
Enable: true, |
|
||||
SearchDepth: 3, |
|
||||
RealTimeSearch: true, |
|
||||
ResultMaxTokens: 500, |
|
||||
}, |
|
||||
Temperature: 0.5, |
|
||||
} |
|
||||
|
|
||||
ctx, cancel := context.WithTimeout(context.Background(), 60*time.Second) |
|
||||
defer cancel() |
|
||||
|
|
||||
response, err := client.ChatWithSearch(ctx, request) |
|
||||
if err != nil { |
|
||||
fmt.Printf("AI回答错误:%v", err) |
|
||||
return |
|
||||
} |
|
||||
|
|
||||
// 处理响应
|
|
||||
if len(response.Choices) > 0 { |
|
||||
choice := response.Choices[0] |
|
||||
fmt.Println("AI回答:", choice.Message.Content) |
|
||||
|
|
||||
if len(choice.SearchResult) > 0 { |
|
||||
fmt.Println("\n引用的搜索结果:") |
|
||||
for _, result := range choice.SearchResult { |
|
||||
fmt.Printf("标题: %s\n摘要: %s\n链接: %s\n\n", |
|
||||
result.Title, |
|
||||
ellipsis(result.Snippet, 100), |
|
||||
result.URL) |
|
||||
} |
|
||||
} |
|
||||
} |
|
||||
} |
|
||||
|
|
||||
// 辅助函数:截断长文本
|
|
||||
func ellipsis(s string, maxLen int) string { |
|
||||
if len(s) <= maxLen { |
|
||||
return s |
|
||||
} |
|
||||
return s[:maxLen] + "..." |
|
||||
} |
|
||||
|
|
||||
func Test_Sys(t *testing.T) { |
|
||||
if sys, err := deepseek.NewSys( |
|
||||
deepseek.SetAppkey("sk-3adfd188a3134e718bbf704f525aff17"), |
|
||||
); err != nil { |
|
||||
fmt.Printf("Sys Init err:%v", err) |
|
||||
} else { |
|
||||
result, err := sys.Chat([]deepseek.Message{{ |
|
||||
Role: "user", |
|
||||
Content: "茅台今日股价分析需包含技术指标", |
|
||||
Web: true, |
|
||||
}}) |
|
||||
fmt.Printf(" result:%v err:%v", result, err) |
|
||||
} |
|
||||
} |
|
||||
@ -1,196 +0,0 @@ |
|||||
package dify |
|
||||
|
|
||||
import "encoding/json" |
|
||||
|
|
||||
const ( |
|
||||
baseurl = "https://api.dify.ai/v1/chat-messages" |
|
||||
workflowsurl = "https://api.dify.ai/v1/workflows/run" |
|
||||
) |
|
||||
|
|
||||
type ( |
|
||||
ISys interface { |
|
||||
ChatForChan(msg string, result chan string) (err error) |
|
||||
Workflows(msg string, result chan string) (err error) |
|
||||
} |
|
||||
|
|
||||
File struct { |
|
||||
Type string `json:"type"` |
|
||||
TransferMethod string `json:"transfer_method"` |
|
||||
URL string `json:"url"` |
|
||||
} |
|
||||
|
|
||||
ChatMessageRequest struct { |
|
||||
Inputs map[string]interface{} `json:"inputs"` |
|
||||
Query string `json:"query"` |
|
||||
ResponseMode string `json:"response_mode"` |
|
||||
ConversationID string `json:"conversation_id"` |
|
||||
User string `json:"user"` |
|
||||
Files []File `json:"files"` |
|
||||
} |
|
||||
|
|
||||
BaseEvent struct { |
|
||||
Event string `json:"event"` |
|
||||
ConversationID string `json:"conversation_id"` |
|
||||
MessageID string `json:"message_id,omitempty"` |
|
||||
CreatedAt int64 `json:"created_at"` |
|
||||
TaskID string `json:"task_id,omitempty"` // 仅 TTS 相关事件存在
|
|
||||
} |
|
||||
|
|
||||
MessageEvent struct { |
|
||||
BaseEvent |
|
||||
Answer string `json:"answer"` |
|
||||
} |
|
||||
|
|
||||
Usage struct { |
|
||||
PromptTokens int `json:"prompt_tokens"` |
|
||||
PromptUnitPrice string `json:"prompt_unit_price"` |
|
||||
PromptPrice string `json:"prompt_price"` |
|
||||
CompletionTokens int `json:"completion_tokens"` |
|
||||
CompletionUnitPrice string `json:"completion_unit_price"` |
|
||||
CompletionPrice string `json:"completion_price"` |
|
||||
TotalTokens int `json:"total_tokens"` |
|
||||
TotalPrice string `json:"total_price"` |
|
||||
Currency string `json:"currency"` |
|
||||
Latency float64 `json:"latency"` |
|
||||
} |
|
||||
|
|
||||
RetrieverResource struct { |
|
||||
Position int `json:"position"` |
|
||||
DatasetID string `json:"dataset_id"` |
|
||||
DatasetName string `json:"dataset_name"` |
|
||||
DocumentID string `json:"document_id"` |
|
||||
DocumentName string `json:"document_name"` |
|
||||
SegmentID string `json:"segment_id"` |
|
||||
Score float64 `json:"score"` |
|
||||
Content string `json:"content"` |
|
||||
} |
|
||||
|
|
||||
MessageEndEvent struct { |
|
||||
BaseEvent |
|
||||
ID string `json:"id"` |
|
||||
Metadata struct { |
|
||||
Usage Usage `json:"usage"` |
|
||||
RetrieverResources []RetrieverResource `json:"retriever_resources"` |
|
||||
} `json:"metadata"` |
|
||||
} |
|
||||
|
|
||||
TTSMessageEvent struct { |
|
||||
BaseEvent |
|
||||
Audio string `json:"audio"` // Base64 编码的音频数据
|
|
||||
} |
|
||||
|
|
||||
TTSMessageEndEvent struct { |
|
||||
BaseEvent |
|
||||
Audio string `json:"audio"` // 可能为空或包含最终音频
|
|
||||
} |
|
||||
|
|
||||
WorkflowResponse struct { |
|
||||
TaskID string `json:"task_id"` |
|
||||
WorkflowRunID string `json:"workflow_run_id"` |
|
||||
Data struct { |
|
||||
ID string `json:"id"` |
|
||||
WorkflowID string `json:"workflow_id"` |
|
||||
Status string `json:"status"` |
|
||||
Outputs struct { |
|
||||
Text string `json:"text"` |
|
||||
} `json:"outputs"` |
|
||||
Error interface{} `json:"error"` // 使用 interface{} 因为可能是 null 或其他类型
|
|
||||
ElapsedTime float64 `json:"elapsed_time"` |
|
||||
TotalTokens int `json:"total_tokens"` |
|
||||
TotalSteps int `json:"total_steps"` |
|
||||
CreatedAt int64 `json:"created_at"` |
|
||||
FinishedAt int64 `json:"finished_at"` |
|
||||
} `json:"data"` |
|
||||
} |
|
||||
|
|
||||
// 通用事件结构
|
|
||||
EventData struct { |
|
||||
Event string `json:"event"` |
|
||||
TaskID string `json:"task_id"` |
|
||||
Data json.RawMessage `json:"data"` // 这里使用 RawMessage 延迟解析
|
|
||||
} |
|
||||
|
|
||||
// workflow_started 事件结构
|
|
||||
WorkflowStartedData struct { |
|
||||
WorkflowRunID string `json:"workflow_run_id"` |
|
||||
ID string `json:"id"` |
|
||||
WorkflowID string `json:"workflow_id"` |
|
||||
SequenceNum int `json:"sequence_number"` |
|
||||
CreatedAt int64 `json:"created_at"` |
|
||||
} |
|
||||
|
|
||||
// node_started 事件结构
|
|
||||
NodeStartedData struct { |
|
||||
WorkflowRunID string `json:"workflow_run_id"` |
|
||||
ID string `json:"id"` |
|
||||
NodeID string `json:"node_id"` |
|
||||
NodeType string `json:"node_type"` |
|
||||
Title string `json:"title"` |
|
||||
Index int `json:"index"` |
|
||||
Inputs map[string]interface{} `json:"inputs"` |
|
||||
CreatedAt int64 `json:"created_at"` |
|
||||
} |
|
||||
|
|
||||
// node_finished 事件结构
|
|
||||
NodeFinishedData struct { |
|
||||
WorkflowRunID string `json:"workflow_run_id"` |
|
||||
ID string `json:"id"` |
|
||||
NodeID string `json:"node_id"` |
|
||||
NodeType string `json:"node_type"` |
|
||||
Title string `json:"title"` |
|
||||
Index int `json:"index"` |
|
||||
Inputs map[string]interface{} `json:"inputs"` |
|
||||
Outputs map[string]interface{} `json:"outputs"` |
|
||||
Status string `json:"status"` |
|
||||
ElapsedTime float64 `json:"elapsed_time"` |
|
||||
ExecutionMeta struct { |
|
||||
TotalTokens int `json:"total_tokens"` |
|
||||
TotalPrice float64 `json:"total_price"` |
|
||||
Currency string `json:"currency"` |
|
||||
} `json:"execution_metadata"` |
|
||||
CreatedAt int64 `json:"created_at"` |
|
||||
} |
|
||||
|
|
||||
// workflow_finished 事件结构
|
|
||||
WorkflowFinishedData struct { |
|
||||
WorkflowRunID string `json:"workflow_run_id"` |
|
||||
ID string `json:"id"` |
|
||||
WorkflowID string `json:"workflow_id"` |
|
||||
Outputs map[string]interface{} `json:"outputs"` |
|
||||
Status string `json:"status"` |
|
||||
ElapsedTime float64 `json:"elapsed_time"` |
|
||||
TotalTokens int `json:"total_tokens"` |
|
||||
TotalSteps string `json:"total_steps"` |
|
||||
CreatedAt int64 `json:"created_at"` |
|
||||
FinishedAt int64 `json:"finished_at"` |
|
||||
} |
|
||||
|
|
||||
// tts_message 事件结构
|
|
||||
TTSMessageData struct { |
|
||||
ConversationID string `json:"conversation_id"` |
|
||||
MessageID string `json:"message_id"` |
|
||||
CreatedAt int64 `json:"created_at"` |
|
||||
TaskID string `json:"task_id"` |
|
||||
Audio string `json:"audio"` |
|
||||
} |
|
||||
) |
|
||||
|
|
||||
var defsys ISys |
|
||||
|
|
||||
func OnInit(config map[string]interface{}, option ...Option) (err error) { |
|
||||
defsys, err = newSys(newOptions(config, option...)) |
|
||||
return |
|
||||
} |
|
||||
|
|
||||
func NewSys(option ...Option) (sys ISys, err error) { |
|
||||
sys, err = newSys(newOptionsByOption(option...)) |
|
||||
return |
|
||||
} |
|
||||
|
|
||||
func ChatForChan(message string, result chan string) (err error) { |
|
||||
return defsys.ChatForChan(message, result) |
|
||||
} |
|
||||
|
|
||||
func Workflows(msg string, result chan string) (err error) { |
|
||||
return defsys.Workflows(msg, result) |
|
||||
} |
|
||||
@ -1,203 +0,0 @@ |
|||||
package dify |
|
||||
|
|
||||
import ( |
|
||||
"bufio" |
|
||||
"bytes" |
|
||||
"context" |
|
||||
"encoding/json" |
|
||||
"fmt" |
|
||||
"log" |
|
||||
"net/http" |
|
||||
"strings" |
|
||||
"time" |
|
||||
) |
|
||||
|
|
||||
func newSys(options Options) (sys *Dify, err error) { |
|
||||
sys = &Dify{ |
|
||||
options: options, |
|
||||
} |
|
||||
return |
|
||||
} |
|
||||
|
|
||||
type Dify struct { |
|
||||
options Options |
|
||||
} |
|
||||
|
|
||||
func (this *Dify) ChatForChan(msg string, result chan string) (err error) { |
|
||||
var ( |
|
||||
body []byte |
|
||||
line []byte |
|
||||
value string |
|
||||
) |
|
||||
// 创建带超时的 context(建议 5-10 分钟,根据实际需求调整)
|
|
||||
ctx, cancel := context.WithTimeout(context.Background(), 5*time.Minute) |
|
||||
defer cancel() |
|
||||
|
|
||||
requestBody := ChatMessageRequest{ |
|
||||
Inputs: make(map[string]interface{}), |
|
||||
Query: msg, |
|
||||
ResponseMode: "streaming", |
|
||||
ConversationID: "", |
|
||||
User: "abc-123", |
|
||||
// Files: []File{
|
|
||||
// {
|
|
||||
// Type: "image",
|
|
||||
// TransferMethod: "remote_url",
|
|
||||
// URL: "https://cloud.dify.ai/logo/logo-site.png",
|
|
||||
// },
|
|
||||
// },
|
|
||||
} |
|
||||
|
|
||||
body, err = json.Marshal(requestBody) |
|
||||
if err != nil { |
|
||||
return |
|
||||
} |
|
||||
|
|
||||
req, _ := http.NewRequestWithContext(ctx, "POST", baseurl, bytes.NewReader(body)) |
|
||||
req.Header.Set("Authorization", "Bearer "+this.options.ApiKey) |
|
||||
req.Header.Set("Content-Type", "application/json") |
|
||||
|
|
||||
client := &http.Client{ |
|
||||
Timeout: 0, // 禁用客户端超时,使用 context 控制
|
|
||||
} |
|
||||
|
|
||||
resp, err := client.Do(req) |
|
||||
if err != nil { |
|
||||
return |
|
||||
} |
|
||||
defer resp.Body.Close() |
|
||||
|
|
||||
if resp.StatusCode != http.StatusOK { |
|
||||
err = fmt.Errorf("code == %d", resp.StatusCode) |
|
||||
} |
|
||||
|
|
||||
// 流式读取核心逻辑
|
|
||||
reader := bufio.NewReader(resp.Body) |
|
||||
for { |
|
||||
select { |
|
||||
case <-ctx.Done(): |
|
||||
log.Println("上下文超时或取消") |
|
||||
return |
|
||||
default: |
|
||||
// 读取直到遇到 \n\n(SSE 事件分隔符)
|
|
||||
line, err = reader.ReadBytes('\n') |
|
||||
if err != nil { |
|
||||
if err.Error() == "EOF" { |
|
||||
close(result) |
|
||||
return |
|
||||
} |
|
||||
return |
|
||||
} |
|
||||
|
|
||||
// 清理数据行
|
|
||||
line = bytes.TrimSpace(line) |
|
||||
if len(line) == 0 { |
|
||||
continue |
|
||||
} |
|
||||
|
|
||||
// 处理 SSE 格式
|
|
||||
if bytes.HasPrefix(line, []byte("data: ")) { |
|
||||
data := bytes.TrimPrefix(line, []byte("data: ")) |
|
||||
if value, err = handleEvent(data); err == nil { |
|
||||
result <- value |
|
||||
} |
|
||||
} |
|
||||
} |
|
||||
} |
|
||||
} |
|
||||
|
|
||||
func handleEvent(data []byte) (result string, err error) { |
|
||||
var base BaseEvent |
|
||||
if err = json.Unmarshal(data, &base); err != nil { |
|
||||
// fmt.Printf("Error parsing base event: %v", err)
|
|
||||
return |
|
||||
} |
|
||||
|
|
||||
switch base.Event { |
|
||||
case "message": |
|
||||
var msg MessageEvent |
|
||||
if err = json.Unmarshal(data, &msg); err != nil { |
|
||||
// fmt.Printf("Error parsing message event: %v", err)
|
|
||||
return |
|
||||
} |
|
||||
result = msg.Answer |
|
||||
// fmt.Printf("[MSG] %s", msg.Answer)
|
|
||||
case "message_end": |
|
||||
var end MessageEndEvent |
|
||||
if err = json.Unmarshal(data, &end); err != nil { |
|
||||
// fmt.Printf("Error parsing message_end event: %v", err)
|
|
||||
return |
|
||||
} |
|
||||
// fmt.Printf("\n[END] Usage: %+v", end.Metadata.Usage)
|
|
||||
|
|
||||
case "tts_message": |
|
||||
var tts TTSMessageEvent |
|
||||
if err = json.Unmarshal(data, &tts); err != nil { |
|
||||
// fmt.Printf("Error parsing tts_message event: %v", err)
|
|
||||
return |
|
||||
} |
|
||||
// 处理音频数据(base64 解码等)
|
|
||||
// fmt.Printf("[TTS] Received audio chunk (%d bytes)\n", len(tts.Audio))
|
|
||||
|
|
||||
case "tts_message_end": |
|
||||
var ttsEnd TTSMessageEndEvent |
|
||||
if err = json.Unmarshal(data, &ttsEnd); err != nil { |
|
||||
// fmt.Printf("Error parsing tts_message_end event: %v", err)
|
|
||||
return |
|
||||
} |
|
||||
// fmt.Println("[TTS END] Audio stream completed")
|
|
||||
default: |
|
||||
// fmt.Printf("Unknown event type: %s", base.Event)
|
|
||||
} |
|
||||
return |
|
||||
} |
|
||||
|
|
||||
func (this *Dify) Workflows(msg string, result chan string) (err error) { |
|
||||
var ( |
|
||||
jsonData []byte |
|
||||
req *http.Request |
|
||||
) |
|
||||
defer close(result) |
|
||||
// 请求数据
|
|
||||
requestData := map[string]interface{}{ |
|
||||
"inputs": map[string]interface{}{ |
|
||||
"user_input": msg, |
|
||||
}, // 工作流输入参数
|
|
||||
"response_mode": "streaming", // 响应模式
|
|
||||
"user": "abc-123", // 用户标识
|
|
||||
} |
|
||||
|
|
||||
// 将请求数据转换为 JSON
|
|
||||
jsonData, err = json.Marshal(requestData) |
|
||||
if err != nil { |
|
||||
return |
|
||||
} |
|
||||
|
|
||||
// 创建 HTTP 请求
|
|
||||
req, err = http.NewRequest("POST", workflowsurl, strings.NewReader(string(jsonData))) |
|
||||
if err != nil { |
|
||||
return |
|
||||
} |
|
||||
req.Header.Set("Authorization", "Bearer app-Dh6RBY4u4G8Kfk4l9sle1UPl") |
|
||||
req.Header.Set("Content-Type", "application/json") |
|
||||
|
|
||||
// 发送请求
|
|
||||
client := &http.Client{} |
|
||||
resp, err := client.Do(req) |
|
||||
if err != nil { |
|
||||
return |
|
||||
} |
|
||||
defer resp.Body.Close() |
|
||||
scanner := bufio.NewScanner(resp.Body) |
|
||||
for scanner.Scan() { |
|
||||
line := scanner.Text() |
|
||||
if strings.HasPrefix(line, "data: ") { |
|
||||
line = strings.TrimPrefix(line, "data: ") |
|
||||
result <- line |
|
||||
} |
|
||||
} |
|
||||
if err = scanner.Err(); err != nil { |
|
||||
return |
|
||||
} |
|
||||
return |
|
||||
} |
|
||||
@ -1,44 +0,0 @@ |
|||||
package dify |
|
||||
|
|
||||
import ( |
|
||||
"yunyan/lego/sys/log" |
|
||||
"yunyan/lego/utils/mapstructure" |
|
||||
) |
|
||||
|
|
||||
type Option func(*Options) |
|
||||
type Options struct { |
|
||||
Debug bool //日志是否开启
|
|
||||
Log log.ILogger |
|
||||
ApiKey string |
|
||||
} |
|
||||
|
|
||||
func SetApiKey(v string) Option { |
|
||||
return func(o *Options) { |
|
||||
o.ApiKey = v |
|
||||
} |
|
||||
} |
|
||||
|
|
||||
func newOptions(config map[string]interface{}, opts ...Option) Options { |
|
||||
options := Options{} |
|
||||
if config != nil { |
|
||||
mapstructure.Decode(config, &options) |
|
||||
} |
|
||||
for _, o := range opts { |
|
||||
o(&options) |
|
||||
} |
|
||||
if options.Log == nil { |
|
||||
options.Log = log.NewTurnlog(options.Debug, log.Clone("sys.dify", 3)) |
|
||||
} |
|
||||
return options |
|
||||
} |
|
||||
|
|
||||
func newOptionsByOption(opts ...Option) Options { |
|
||||
options := Options{} |
|
||||
for _, o := range opts { |
|
||||
o(&options) |
|
||||
} |
|
||||
if options.Log == nil { |
|
||||
options.Log = log.NewTurnlog(options.Debug, log.Clone("sys.dify", 3)) |
|
||||
} |
|
||||
return options |
|
||||
} |
|
||||
@ -1,160 +0,0 @@ |
|||||
package dify_test |
|
||||
|
|
||||
import ( |
|
||||
"bufio" |
|
||||
"yunyan/sys/dify" |
|
||||
"encoding/json" |
|
||||
"fmt" |
|
||||
"net/http" |
|
||||
"strings" |
|
||||
"testing" |
|
||||
|
|
||||
"github.com/go-resty/resty/v2" |
|
||||
) |
|
||||
|
|
||||
func Test_Sys_Chat(t *testing.T) { |
|
||||
if sys, err := dify.NewSys( |
|
||||
dify.SetApiKey("app-7wHwovYZsZBveMPpi1k1ZN8c"), |
|
||||
); err != nil { |
|
||||
fmt.Printf("Sys Init err:%v", err) |
|
||||
} else { |
|
||||
result := make(chan string) |
|
||||
go func() { |
|
||||
err := sys.ChatForChan("你好! 你的名字?", result) |
|
||||
if err != nil { |
|
||||
fmt.Printf(" err:%v", err) |
|
||||
return |
|
||||
} |
|
||||
}() |
|
||||
|
|
||||
for v := range result { |
|
||||
fmt.Printf(" result:%v", v) |
|
||||
} |
|
||||
} |
|
||||
} |
|
||||
|
|
||||
func Test_Sys_Workflows(t *testing.T) { |
|
||||
|
|
||||
client := resty.New() |
|
||||
// Dify API 端点
|
|
||||
url := "https://api.dify.ai/v1/workflows/run" |
|
||||
// 请求数据
|
|
||||
requestData := map[string]interface{}{ |
|
||||
"inputs": map[string]interface{}{ |
|
||||
"user_input": "武汉的天气如何?", |
|
||||
}, // 工作流输入参数
|
|
||||
"response_mode": "blocking", // 响应模式
|
|
||||
"user": "abc-123", // 用户标识
|
|
||||
} |
|
||||
|
|
||||
// 发送 POST 请求
|
|
||||
resp, err := client.R(). |
|
||||
SetHeader("Authorization", "Bearer app-Dh6RBY4u4G8Kfk4l9sle1UPl"). |
|
||||
SetHeader("Content-Type", "application/json"). |
|
||||
SetBody(requestData). |
|
||||
Post(url) |
|
||||
|
|
||||
if err != nil { |
|
||||
fmt.Println("Error sending request:", err) |
|
||||
return |
|
||||
} |
|
||||
|
|
||||
// 打印响应
|
|
||||
fmt.Println("Response Status:", resp.Status()) |
|
||||
// fmt.Println("Response Body:", resp.String())
|
|
||||
wresp := &dify.WorkflowResponse{} |
|
||||
err = json.Unmarshal(resp.Body(), wresp) |
|
||||
fmt.Printf("Response:%+v", wresp) |
|
||||
} |
|
||||
|
|
||||
func Test_Sys_WorkflowsForSatem(t *testing.T) { |
|
||||
// Dify API 端点
|
|
||||
url := "https://api.dify.ai/v1/workflows/run" |
|
||||
// 请求数据
|
|
||||
requestData := map[string]interface{}{ |
|
||||
"inputs": map[string]interface{}{ |
|
||||
"user_input": "武汉的天气如何?", |
|
||||
}, // 工作流输入参数
|
|
||||
"response_mode": "streaming", // 响应模式
|
|
||||
"user": "abc-123", // 用户标识
|
|
||||
} |
|
||||
|
|
||||
// 将请求数据转换为 JSON
|
|
||||
jsonData, err := json.Marshal(requestData) |
|
||||
if err != nil { |
|
||||
t.Fatalf("Error marshalling JSON: %v", err) |
|
||||
} |
|
||||
|
|
||||
// 创建 HTTP 请求
|
|
||||
req, err := http.NewRequest("POST", url, strings.NewReader(string(jsonData))) |
|
||||
if err != nil { |
|
||||
t.Fatalf("Error creating request: %v", err) |
|
||||
} |
|
||||
req.Header.Set("Authorization", "Bearer app-Dh6RBY4u4G8Kfk4l9sle1UPl") |
|
||||
req.Header.Set("Content-Type", "application/json") |
|
||||
|
|
||||
// 发送请求
|
|
||||
client := &http.Client{} |
|
||||
resp, err := client.Do(req) |
|
||||
if err != nil { |
|
||||
t.Fatalf("Error sending request: %v", err) |
|
||||
} |
|
||||
defer resp.Body.Close() |
|
||||
|
|
||||
defer resp.Body.Close() |
|
||||
|
|
||||
scanner := bufio.NewScanner(resp.Body) |
|
||||
for scanner.Scan() { |
|
||||
line := scanner.Text() |
|
||||
if strings.HasPrefix(line, "data: ") { |
|
||||
line = strings.TrimPrefix(line, "data: ") |
|
||||
var event dify.EventData |
|
||||
if err := json.Unmarshal([]byte(line), &event); err != nil { |
|
||||
fmt.Printf("Error parsing JSON: %v\n", err) |
|
||||
continue |
|
||||
} |
|
||||
parseEvent(event) // 解析不同类型的事件
|
|
||||
} |
|
||||
} |
|
||||
|
|
||||
if err := scanner.Err(); err != nil { |
|
||||
t.Fatalf("Error reading stream response: %v", err) |
|
||||
} |
|
||||
} |
|
||||
|
|
||||
// 解析数据
|
|
||||
func parseEvent(event dify.EventData) { |
|
||||
switch event.Event { |
|
||||
case "workflow_started": |
|
||||
var data dify.WorkflowStartedData |
|
||||
if err := json.Unmarshal(event.Data, &data); err == nil { |
|
||||
fmt.Printf("Workflow started: %+v\n", data) |
|
||||
} |
|
||||
|
|
||||
case "node_started": |
|
||||
var data dify.NodeStartedData |
|
||||
if err := json.Unmarshal(event.Data, &data); err == nil { |
|
||||
fmt.Printf("Node started: %+v\n", data) |
|
||||
} |
|
||||
|
|
||||
case "node_finished": |
|
||||
var data dify.NodeFinishedData |
|
||||
if err := json.Unmarshal(event.Data, &data); err == nil { |
|
||||
fmt.Printf("Node finished: %+v\n", data) |
|
||||
} |
|
||||
|
|
||||
case "workflow_finished": |
|
||||
var data dify.WorkflowFinishedData |
|
||||
if err := json.Unmarshal(event.Data, &data); err == nil { |
|
||||
fmt.Printf("Workflow finished: %+v\n", data) |
|
||||
} |
|
||||
|
|
||||
case "tts_message": |
|
||||
var data dify.TTSMessageData |
|
||||
if err := json.Unmarshal(event.Data, &data); err == nil { |
|
||||
fmt.Printf("TTS message: %+v\n", data) |
|
||||
} |
|
||||
default: |
|
||||
fmt.Printf("Unknown event: %s\n", event.Event) |
|
||||
} |
|
||||
} |
|
||||
Some files were not shown because too many files changed in this diff
Loading…
Reference in new issue