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