Init
This commit is contained in:
@@ -0,0 +1,117 @@
|
||||
package wxwork
|
||||
|
||||
import (
|
||||
"log/slog"
|
||||
"strings"
|
||||
"sync"
|
||||
|
||||
"github.com/panjf2000/ants/v2"
|
||||
"github.com/silenceper/wechat/v2/work/kf"
|
||||
)
|
||||
|
||||
const (
|
||||
callbackWorkerCount = 16 // worker 数量
|
||||
callbackMaxBlockingTasks = 10000 // 队列大小
|
||||
callbackMsgTypeEvent = "event"
|
||||
callbackEventKFMsgOrEvent = "kf_msg_or_event"
|
||||
)
|
||||
|
||||
type CallbackHandler func(message kf.CallbackMessage)
|
||||
|
||||
var (
|
||||
callbackHandlerMu sync.RWMutex
|
||||
callbackHandlers map[string]CallbackHandler
|
||||
callbackPool *ants.PoolWithFunc
|
||||
)
|
||||
|
||||
func init() {
|
||||
var err error
|
||||
|
||||
callbackHandlers = make(map[string]CallbackHandler)
|
||||
callbackPool, err = ants.NewPoolWithFunc(
|
||||
callbackWorkerCount,
|
||||
dispatchCallbackMessage,
|
||||
ants.WithMaxBlockingTasks(callbackMaxBlockingTasks),
|
||||
)
|
||||
if err != nil {
|
||||
slog.Error("init wxwork callback pool failed", "error", err)
|
||||
}
|
||||
}
|
||||
|
||||
// ConsumeCallback 异步消费企业微信客服回调。
|
||||
func ConsumeCallback(message kf.CallbackMessage) error {
|
||||
if callbackPool == nil {
|
||||
slog.Error("wxwork callback pool is not initialized",
|
||||
"msg_type", message.MsgType,
|
||||
"event", message.Event,
|
||||
"open_kfid", message.OpenKfID,
|
||||
)
|
||||
return nil
|
||||
}
|
||||
|
||||
if err := callbackPool.Invoke(message); err != nil {
|
||||
slog.Error("invoke wxwork callback failed",
|
||||
"msg_type", message.MsgType,
|
||||
"event", message.Event,
|
||||
"open_kfid", message.OpenKfID,
|
||||
"error", err,
|
||||
)
|
||||
return err
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// RegHandler 注册企业微信客服回调处理器。
|
||||
func RegHandler(msgType, event string, handler CallbackHandler) {
|
||||
callbackHandlerMu.Lock()
|
||||
defer callbackHandlerMu.Unlock()
|
||||
|
||||
key := buildCallbackKey(msgType, event)
|
||||
if callbackHandlers[key] != nil {
|
||||
slog.Error("duplicate wxwork callback handler registration", "key", key)
|
||||
return
|
||||
}
|
||||
callbackHandlers[key] = handler
|
||||
}
|
||||
|
||||
// buildCallbackKey 生成客服回调处理器的注册键。
|
||||
func buildCallbackKey(msgType, event string) string {
|
||||
msgType = strings.TrimSpace(msgType)
|
||||
event = strings.TrimSpace(event)
|
||||
if msgType == callbackMsgTypeEvent {
|
||||
return msgType + ":" + event
|
||||
}
|
||||
return msgType
|
||||
}
|
||||
|
||||
func dispatchCallbackMessage(value any) {
|
||||
message, ok := value.(kf.CallbackMessage)
|
||||
if !ok {
|
||||
slog.Error("invalid wxwork callback message type")
|
||||
return
|
||||
}
|
||||
|
||||
handler := getCallbackHandler(message)
|
||||
if handler == nil {
|
||||
return
|
||||
}
|
||||
handler(message)
|
||||
}
|
||||
|
||||
func getCallbackHandler(message kf.CallbackMessage) CallbackHandler {
|
||||
callbackHandlerMu.RLock()
|
||||
defer callbackHandlerMu.RUnlock()
|
||||
|
||||
key := buildCallbackKey(message.MsgType, message.Event)
|
||||
handler, ok := callbackHandlers[key]
|
||||
if ok {
|
||||
return handler
|
||||
}
|
||||
|
||||
slog.Warn("wxwork callback handler not found",
|
||||
"key", key,
|
||||
"msg_type", message.MsgType,
|
||||
"event", message.Event,
|
||||
)
|
||||
return nil
|
||||
}
|
||||
@@ -0,0 +1,255 @@
|
||||
package wxwork
|
||||
|
||||
import (
|
||||
"crypto/hmac"
|
||||
"crypto/rand"
|
||||
"crypto/sha256"
|
||||
"cs-agent/internal/pkg/dto/response"
|
||||
"cs-agent/internal/pkg/errorsx"
|
||||
"encoding/base64"
|
||||
"encoding/hex"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"net/url"
|
||||
"strings"
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
"github.com/silenceper/wechat/v2/work/oauth"
|
||||
)
|
||||
|
||||
const (
|
||||
StateTTL = 5 * time.Minute
|
||||
LoginTicketTTL = 1 * time.Minute
|
||||
defaultLoginNextPath = "/dashboard"
|
||||
)
|
||||
|
||||
var (
|
||||
errStateInvalid = fmt.Errorf("企业微信登录状态无效")
|
||||
loginTicketStore sync.Map
|
||||
)
|
||||
|
||||
type statePayload struct {
|
||||
Next string `json:"next"`
|
||||
Nonce string `json:"nonce"`
|
||||
ExpiredAt int64 `json:"expiredAt"`
|
||||
}
|
||||
|
||||
type loginTicket struct {
|
||||
Response *response.LoginResponse
|
||||
ExpiredAt time.Time
|
||||
}
|
||||
|
||||
func BuildLoginURL(state string) (string, error) {
|
||||
if !Enabled() {
|
||||
return "", fmt.Errorf("企业微信登录未启用")
|
||||
}
|
||||
if strings.TrimSpace(wxCfg.OAuthRedirect) == "" {
|
||||
return "", fmt.Errorf("企业微信登录回调地址未配置")
|
||||
}
|
||||
if strings.TrimSpace(wxCfg.AgentID) == "" {
|
||||
return "", fmt.Errorf("企业微信 AgentID 未配置")
|
||||
}
|
||||
return fmt.Sprintf(
|
||||
"https://open.weixin.qq.com/connect/oauth2/authorize?appid=%s&redirect_uri=%s&response_type=code&scope=snsapi_privateinfo&agentid=%s&state=%s#wechat_redirect",
|
||||
url.QueryEscape(strings.TrimSpace(wxCfg.CorpID)),
|
||||
url.QueryEscape(strings.TrimSpace(wxCfg.OAuthRedirect)),
|
||||
url.QueryEscape(strings.TrimSpace(wxCfg.AgentID)),
|
||||
url.QueryEscape(strings.TrimSpace(state)),
|
||||
), nil
|
||||
}
|
||||
|
||||
func BuildQRCodeLoginURL(state string) (string, error) {
|
||||
if !Enabled() {
|
||||
return "", fmt.Errorf("企业微信登录未启用")
|
||||
}
|
||||
if strings.TrimSpace(wxCfg.OAuthRedirect) == "" {
|
||||
return "", fmt.Errorf("企业微信登录回调地址未配置")
|
||||
}
|
||||
if strings.TrimSpace(wxCfg.AgentID) == "" {
|
||||
return "", fmt.Errorf("企业微信 AgentID 未配置")
|
||||
}
|
||||
return fmt.Sprintf(
|
||||
"https://open.work.weixin.qq.com/wwopen/sso/qrConnect?appid=%s&agentid=%s&redirect_uri=%s&state=%s",
|
||||
url.QueryEscape(strings.TrimSpace(wxCfg.CorpID)),
|
||||
url.QueryEscape(strings.TrimSpace(wxCfg.AgentID)),
|
||||
url.QueryEscape(strings.TrimSpace(wxCfg.OAuthRedirect)),
|
||||
url.QueryEscape(strings.TrimSpace(state)),
|
||||
), nil
|
||||
}
|
||||
|
||||
func CreateState(next string) (string, error) {
|
||||
secret := strings.TrimSpace(StateSecret())
|
||||
if secret == "" {
|
||||
return "", errorsx.BusinessError(1, "企业微信登录密钥未配置")
|
||||
}
|
||||
payload := statePayload{
|
||||
Next: sanitizeNextPath(next),
|
||||
ExpiredAt: time.Now().Add(StateTTL).Unix(),
|
||||
}
|
||||
nonce, err := randomToken("ws_")
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
payload.Nonce = nonce
|
||||
|
||||
body, err := json.Marshal(payload)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
encoded := base64.RawURLEncoding.EncodeToString(body)
|
||||
return encoded + "." + signState(encoded, secret), nil
|
||||
}
|
||||
|
||||
func ParseState(state string) (string, error) {
|
||||
secret := strings.TrimSpace(StateSecret())
|
||||
if secret == "" {
|
||||
return "", errStateInvalid
|
||||
}
|
||||
parts := strings.Split(strings.TrimSpace(state), ".")
|
||||
if len(parts) != 2 {
|
||||
return "", errStateInvalid
|
||||
}
|
||||
if !hmac.Equal([]byte(parts[1]), []byte(signState(parts[0], secret))) {
|
||||
return "", errStateInvalid
|
||||
}
|
||||
body, err := base64.RawURLEncoding.DecodeString(parts[0])
|
||||
if err != nil {
|
||||
return "", errStateInvalid
|
||||
}
|
||||
payload := statePayload{}
|
||||
if err = json.Unmarshal(body, &payload); err != nil {
|
||||
return "", errStateInvalid
|
||||
}
|
||||
if payload.ExpiredAt <= time.Now().Unix() {
|
||||
return "", errStateInvalid
|
||||
}
|
||||
return sanitizeNextPath(payload.Next), nil
|
||||
}
|
||||
|
||||
func IssueLoginTicket(loginResp *response.LoginResponse) (string, error) {
|
||||
if loginResp == nil {
|
||||
return "", fmt.Errorf("登录结果不能为空")
|
||||
}
|
||||
ticket, err := randomToken("wlt_")
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
cleanupExpiredLoginTickets()
|
||||
loginTicketStore.Store(ticket, loginTicket{
|
||||
Response: loginResp,
|
||||
ExpiredAt: time.Now().Add(LoginTicketTTL),
|
||||
})
|
||||
return ticket, nil
|
||||
}
|
||||
|
||||
func ConsumeLoginTicket(ticket string) (*response.LoginResponse, error) {
|
||||
ticket = strings.TrimSpace(ticket)
|
||||
if ticket == "" {
|
||||
return nil, errorsx.InvalidParam("ticket 不能为空")
|
||||
}
|
||||
value, ok := loginTicketStore.LoadAndDelete(ticket)
|
||||
if !ok {
|
||||
return nil, errorsx.Unauthorized("登录票据无效或已过期")
|
||||
}
|
||||
record, ok := value.(loginTicket)
|
||||
if !ok || record.Response == nil || time.Now().After(record.ExpiredAt) {
|
||||
return nil, errorsx.Unauthorized("登录票据无效或已过期")
|
||||
}
|
||||
return record.Response, nil
|
||||
}
|
||||
|
||||
func signState(content, secret string) string {
|
||||
mac := hmac.New(sha256.New, []byte(secret))
|
||||
_, _ = mac.Write([]byte(content))
|
||||
return hex.EncodeToString(mac.Sum(nil))
|
||||
}
|
||||
|
||||
func cleanupExpiredLoginTickets() {
|
||||
now := time.Now()
|
||||
loginTicketStore.Range(func(key, value any) bool {
|
||||
record, ok := value.(loginTicket)
|
||||
if !ok || now.After(record.ExpiredAt) {
|
||||
loginTicketStore.Delete(key)
|
||||
}
|
||||
return true
|
||||
})
|
||||
}
|
||||
|
||||
func sanitizeNextPath(next string) string {
|
||||
next = strings.TrimSpace(next)
|
||||
if next == "" || !strings.HasPrefix(next, "/") || strings.HasPrefix(next, "//") {
|
||||
return defaultLoginNextPath
|
||||
}
|
||||
return next
|
||||
}
|
||||
|
||||
func randomToken(prefix string) (string, error) {
|
||||
buf := make([]byte, 24)
|
||||
if _, err := rand.Read(buf); err != nil {
|
||||
return "", err
|
||||
}
|
||||
return prefix + hex.EncodeToString(buf), nil
|
||||
}
|
||||
|
||||
func GetUserDetail(code string) (*LoginUser, error) {
|
||||
if !Enabled() {
|
||||
return nil, fmt.Errorf("企业微信登录未启用")
|
||||
}
|
||||
code = strings.TrimSpace(code)
|
||||
if code == "" {
|
||||
return nil, fmt.Errorf("微信授权 code 不能为空")
|
||||
}
|
||||
|
||||
oauthClient := w.GetOauth()
|
||||
userInfo, err := oauthClient.GetUserInfo(code)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if strings.TrimSpace(userInfo.UserID) == "" {
|
||||
return nil, fmt.Errorf("当前登录身份不是企业内部成员")
|
||||
}
|
||||
|
||||
ret := &LoginUser{
|
||||
CorpID: wxCfg.CorpID,
|
||||
UserID: strings.TrimSpace(userInfo.UserID),
|
||||
OpenID: strings.TrimSpace(userInfo.OpenID),
|
||||
ExternalUserID: strings.TrimSpace(userInfo.ExternalUserID),
|
||||
UserTicket: strings.TrimSpace(userInfo.UserTicket),
|
||||
UserInfo: userInfo,
|
||||
}
|
||||
|
||||
if ret.UserTicket != "" {
|
||||
if detail, detailErr := oauthClient.GetUserDetail(&oauth.GetUserDetailRequest{UserTicket: ret.UserTicket}); detailErr == nil {
|
||||
ret.UserDetail = detail
|
||||
ret.Avatar = strings.TrimSpace(detail.Avatar)
|
||||
ret.Mobile = strings.TrimSpace(detail.Mobile)
|
||||
ret.Email = strings.TrimSpace(detail.Email)
|
||||
ret.BizMail = strings.TrimSpace(detail.BizMail)
|
||||
}
|
||||
}
|
||||
|
||||
if profile, profileErr := w.GetAddressList().UserGet(ret.UserID); profileErr == nil {
|
||||
ret.UserProfile = profile
|
||||
if strings.TrimSpace(profile.Name) != "" {
|
||||
ret.Name = strings.TrimSpace(profile.Name)
|
||||
}
|
||||
if strings.TrimSpace(profile.Avatar) != "" {
|
||||
ret.Avatar = strings.TrimSpace(profile.Avatar)
|
||||
}
|
||||
if ret.Mobile == "" && strings.TrimSpace(profile.Mobile) != "" {
|
||||
ret.Mobile = strings.TrimSpace(profile.Mobile)
|
||||
}
|
||||
if ret.Email == "" && strings.TrimSpace(profile.Email) != "" {
|
||||
ret.Email = strings.TrimSpace(profile.Email)
|
||||
}
|
||||
if ret.BizMail == "" && strings.TrimSpace(profile.BizMail) != "" {
|
||||
ret.BizMail = strings.TrimSpace(profile.BizMail)
|
||||
}
|
||||
}
|
||||
|
||||
if ret.Name == "" {
|
||||
ret.Name = ret.UserID
|
||||
}
|
||||
return ret, nil
|
||||
}
|
||||
@@ -0,0 +1,70 @@
|
||||
package wxwork
|
||||
|
||||
import (
|
||||
"cs-agent/internal/pkg/config"
|
||||
"strings"
|
||||
|
||||
"github.com/silenceper/wechat/v2/cache"
|
||||
"github.com/silenceper/wechat/v2/work"
|
||||
"github.com/silenceper/wechat/v2/work/addresslist"
|
||||
wxconfig "github.com/silenceper/wechat/v2/work/config"
|
||||
"github.com/silenceper/wechat/v2/work/oauth"
|
||||
)
|
||||
|
||||
var (
|
||||
w *work.Work
|
||||
wxCfg config.WxWorkConfig
|
||||
)
|
||||
|
||||
type LoginUser struct {
|
||||
CorpID string `json:"corpId"`
|
||||
UserID string `json:"userId"`
|
||||
OpenID string `json:"openId,omitempty"`
|
||||
ExternalUserID string `json:"externalUserId,omitempty"`
|
||||
UserTicket string `json:"userTicket,omitempty"`
|
||||
Name string `json:"name,omitempty"`
|
||||
Avatar string `json:"avatar,omitempty"`
|
||||
Mobile string `json:"mobile,omitempty"`
|
||||
Email string `json:"email,omitempty"`
|
||||
BizMail string `json:"bizMail,omitempty"`
|
||||
UserInfo *oauth.GetUserInfoResponse `json:"userInfo,omitempty"`
|
||||
UserDetail *oauth.GetUserDetailResponse `json:"userDetail,omitempty"`
|
||||
UserProfile *addresslist.UserGetResponse `json:"userProfile,omitempty"`
|
||||
}
|
||||
|
||||
func Init() {
|
||||
w = nil
|
||||
wxCfg = config.WxWorkConfig{}
|
||||
cfg := config.Current()
|
||||
if !cfg.WxWork.Enabled {
|
||||
return
|
||||
}
|
||||
wxCfg = cfg.WxWork
|
||||
if strings.TrimSpace(wxCfg.CorpID) == "" || strings.TrimSpace(wxCfg.CorpSecret) == "" {
|
||||
return
|
||||
}
|
||||
w = work.NewWork(&wxconfig.Config{
|
||||
CorpID: wxCfg.CorpID,
|
||||
CorpSecret: wxCfg.CorpSecret,
|
||||
AgentID: wxCfg.AgentID,
|
||||
RasPrivateKey: wxCfg.RSAPrivateKey,
|
||||
Token: wxCfg.Token,
|
||||
EncodingAESKey: wxCfg.EncodingAESKey,
|
||||
Cache: cache.NewMemory(),
|
||||
})
|
||||
}
|
||||
|
||||
func Enabled() bool {
|
||||
return w != nil && wxCfg.Enabled
|
||||
}
|
||||
|
||||
func StateSecret() string {
|
||||
if strings.TrimSpace(wxCfg.StateSecret) != "" {
|
||||
return strings.TrimSpace(wxCfg.StateSecret)
|
||||
}
|
||||
return strings.TrimSpace(wxCfg.CorpSecret)
|
||||
}
|
||||
|
||||
func GetWorkCli() *work.Work {
|
||||
return w
|
||||
}
|
||||
Reference in New Issue
Block a user