Compare commits

...
47 changed files with 4761 additions and 30 deletions
+13 -1
View File
@@ -10,7 +10,7 @@ SCHEMA_META_INDEX_OUTPUT ?= artifacts/schema_meta_index.gob
POLICY_ENV = DWS_POLICY_TMPDIR="$(DWS_POLICY_TMPDIR)" GOTMPDIR="$(POLICY_GOTMPDIR)"
GO_SOURCE_LIST = git ls-files -z --cached --others --exclude-standard -- '*.go'
.PHONY: all help build rebuild test test-plan test-auth-legacy-compat lint format-check fmt policy edition-test interface-integrity authoritative-interface-integrity coverage-gate coverage-gate-platform update-interface-baseline reset-interface-baseline schema-compatibility skill-command-integrity skill-context-budget multi-im-skill-chain-integrity cli-smoke mock-mcp-smoke test-schema-agent-examples generate-schema fetch-mcp-metadata generate-schema-catalog package release release-pre release-stable changelog-pre changelog-stable publish-homebrew-formula setup-hooks
.PHONY: all help build check-safechat test-safechat rebuild test test-plan test-auth-legacy-compat lint format-check fmt policy edition-test interface-integrity authoritative-interface-integrity coverage-gate coverage-gate-platform update-interface-baseline reset-interface-baseline schema-compatibility skill-command-integrity skill-context-budget multi-im-skill-chain-integrity cli-smoke mock-mcp-smoke test-schema-agent-examples generate-schema fetch-mcp-metadata generate-schema-catalog package release release-pre release-stable changelog-pre changelog-stable publish-homebrew-formula setup-hooks
all: setup-hooks fmt lint build test rebuild
@@ -18,6 +18,8 @@ help:
@printf "Available targets:\n"
@printf " make build - Build the dws CLI binary\n"
@printf " make test - Run the Go test suite\n"
@printf " make check-safechat - Compile and vet the SafeChat message-crypto backend (needs CGO)\n"
@printf " make test-safechat - Run the message-crypto tests against the SafeChat backend\n"
@printf " make test-plan - Verify CI test and full-suite coverage package plans cover their scopes exactly once\n"
@printf " make test-auth-legacy-compat - Run stable legacy authentication compatibility regressions\n"
@printf " make lint - Run formatting checks, go vet, and staticcheck\n"
@@ -52,6 +54,16 @@ build:
rebuild:
@./scripts/dev/build.sh
# No dws command imports internal/msgcrypto yet, so a tagged CLI build would
# link nothing extra and look identical to the default binary. Gate the package
# itself until a caller wires it in.
check-safechat:
@CGO_ENABLED=1 $(GO) build -tags safechat ./internal/msgcrypto/...
@CGO_ENABLED=1 $(GO) vet -tags safechat ./internal/msgcrypto/...
test-safechat:
@CGO_ENABLED=1 $(GO) test -count=1 -tags safechat ./internal/msgcrypto/...
test:
@DWS_PACKAGE_VERSION="$(DWS_PACKAGE_VERSION)" $(GO) test -count=1 -timeout=10m ./...
+3
View File
@@ -4,6 +4,8 @@ go 1.25.9
replace gitlab.alibaba-inc.com/aes/aem-go-sdk => ./third_party/aem-go-sdk
replace safechat-go-sdk => ./third_party/safechat-go-sdk
require (
github.com/Microsoft/go-winio v0.6.2
github.com/RealAlexandreAI/json-repair v0.0.15
@@ -24,6 +26,7 @@ require (
golang.org/x/crypto v0.49.0
golang.org/x/sys v0.42.0
golang.org/x/text v0.35.0
safechat-go-sdk v0.0.0
)
require (
+4 -3
View File
@@ -113,9 +113,10 @@ const (
ClientIDPath = "/cli/clientId"
// MCP OAuth endpoints (used when clientId is fetched from MCP).
MCPOAuthTokenPath = "/oauth2/getToken"
MCPRefreshTokenPath = "/oauth2/refreshToken"
MCPRevokeTokenPath = "/oauth2/revokeToken"
MCPOAuthTokenPath = "/oauth2/getToken"
MCPRefreshTokenPath = "/oauth2/refreshToken"
MCPRevokeTokenPath = "/oauth2/revokeToken"
MCPVendorAuthCodePath = "/oauth2/vendorAuthCode"
// App-level access token endpoints (for dws api raw calls).
+192
View File
@@ -0,0 +1,192 @@
// Copyright 2026 Alibaba Group
// Licensed under the Apache License, Version 2.0 (the "License");
// you may not use this file except in compliance with the License.
// You may obtain a copy of the License at
//
// http://www.apache.org/licenses/LICENSE-2.0
//
// Unless required by applicable law or agreed to in writing, software
// distributed under the License is distributed on an "AS IS" BASIS,
// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
// See the License for the specific language governing permissions and
// limitations under the License.
package auth
import (
"bytes"
"context"
"encoding/json"
"fmt"
"io"
"net/http"
"strings"
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/pkg/config"
)
const (
headerUserAccessToken = "x-user-access-token"
headerDWSClientID = "x-dws-client-id"
headerDWSCLIVersion = "x-dws-cli-version"
// VendorAuthCode default / documented portal expiresIn, used only when
// the success body omits a positive value. Callers must still prefer
// the response field.
DefaultVendorAuthCodeExpiresIn = 120
VendorAuthCodeParamError = "PARAM_ERROR"
VendorAuthCodeVendorUnsupported = "VENDOR_UNSUPPORTED"
VendorAuthCodeTokenInvalid = "TOKEN_INVALID"
VendorAuthCodeOrgMismatch = "ORG_MISMATCH"
VendorAuthCodeUserNotInOrg = "USER_NOT_IN_ORG"
VendorAuthCodeVendorNotEnabled = "VENDOR_NOT_ENABLED"
VendorAuthCodeRateLimited = "RATE_LIMITED"
VendorAuthCodeInternalError = "INTERNAL_ERROR"
)
// VendorAuthCodeInput is a POST /oauth2/vendorAuthCode call. The body is
// only vendor + corpId; redirectURI and domain must not be sent.
type VendorAuthCodeInput struct {
AccessToken string
ClientID string
CLIVersion string
LoginRegion LoginRegion
Vendor string
CorpID string
HTTPClient *http.Client
// BaseURL overrides MCPBaseURLForLoginRegion. Tests use it; production
// callers leave it empty.
BaseURL string
}
// VendorAuthCodeResult is the success VO from portal.
type VendorAuthCodeResult struct {
AuthCode string
ExpiresIn int
}
// VendorAuthCodeError is a portal business error carried in an HTTP 200
// ServiceResult body (same envelope as /oauth2/getToken).
type VendorAuthCodeError struct {
Code string
Message string
}
func (e *VendorAuthCodeError) Error() string {
if e == nil {
return "vendorAuthCode failed"
}
if strings.TrimSpace(e.Message) != "" {
return fmt.Sprintf("vendorAuthCode %s: %s", e.Code, e.Message)
}
return fmt.Sprintf("vendorAuthCode %s", e.Code)
}
// Retryable reports whether DWS should retry this portal error once.
func (e *VendorAuthCodeError) Retryable() bool {
if e == nil {
return false
}
switch e.Code {
case VendorAuthCodeTokenInvalid, VendorAuthCodeRateLimited, VendorAuthCodeInternalError:
return true
default:
return false
}
}
// FetchVendorAuthCode POSTs {vendor, corpId} to /oauth2/vendorAuthCode.
// HTTP is expected to be 200; errors are read from body.errorCode.
func FetchVendorAuthCode(ctx context.Context, in VendorAuthCodeInput) (*VendorAuthCodeResult, error) {
vendor := strings.ToLower(strings.TrimSpace(in.Vendor))
corpID := strings.TrimSpace(in.CorpID)
token := strings.TrimSpace(in.AccessToken)
clientID := strings.TrimSpace(in.ClientID)
if token == "" || clientID == "" || vendor == "" || corpID == "" {
return nil, &VendorAuthCodeError{
Code: VendorAuthCodeParamError,
Message: "token, clientId, vendor and corpId are required",
}
}
base := strings.TrimRight(strings.TrimSpace(in.BaseURL), "/")
if base == "" {
base = strings.TrimRight(MCPBaseURLForLoginRegion(in.LoginRegion), "/")
}
endpoint := base + MCPVendorAuthCodePath
payload, err := json.Marshal(struct {
Vendor string `json:"vendor"`
CorpID string `json:"corpId"`
}{Vendor: vendor, CorpID: corpID})
if err != nil {
return nil, fmt.Errorf("marshaling vendorAuthCode request: %w", err)
}
req, err := oauthNewRequest(ctx, http.MethodPost, endpoint, bytes.NewReader(payload))
if err != nil {
return nil, fmt.Errorf("creating vendorAuthCode request: %w", err)
}
req.Header.Set("Content-Type", "application/json")
req.Header.Set(headerUserAccessToken, token)
req.Header.Set(headerDWSClientID, clientID)
if ver := strings.TrimSpace(in.CLIVersion); ver != "" {
req.Header.Set(headerDWSCLIVersion, ver)
}
applyEditionEnterpriseCredentialHeaders(req)
client := in.HTTPClient
if client == nil {
client = oauthHTTPClient
}
resp, err := client.Do(req)
if err != nil {
return nil, fmt.Errorf("sending vendorAuthCode request: %w", err)
}
defer resp.Body.Close()
data, readErr := io.ReadAll(io.LimitReader(resp.Body, config.MaxResponseBodySize))
if resp.StatusCode != http.StatusOK {
if readErr != nil {
data = nil
}
return nil, &HTTPStatusError{
StatusCode: resp.StatusCode,
responseBody: truncateBody(data, 200),
}
}
if readErr != nil {
return nil, fmt.Errorf("reading vendorAuthCode response: %w", readErr)
}
return parseVendorAuthCodeResponse(data)
}
func parseVendorAuthCodeResponse(body []byte) (*VendorAuthCodeResult, error) {
var resp struct {
AuthCode string `json:"authCode"`
ExpiresIn int `json:"expiresIn"`
Success *bool `json:"success"`
ErrorCode string `json:"errorCode"`
ErrorMsg string `json:"errorMsg"`
}
if err := json.Unmarshal(body, &resp); err != nil {
return nil, fmt.Errorf("parsing vendorAuthCode response: %w", err)
}
if resp.ErrorCode != "" || resp.ErrorMsg != "" || (resp.Success != nil && !*resp.Success) {
code := strings.TrimSpace(resp.ErrorCode)
if code == "" {
code = VendorAuthCodeInternalError
}
return nil, &VendorAuthCodeError{Code: code, Message: resp.ErrorMsg}
}
authCode := strings.TrimSpace(resp.AuthCode)
if authCode == "" {
return nil, fmt.Errorf("vendorAuthCode response missing authCode")
}
expiresIn := resp.ExpiresIn
if expiresIn <= 0 {
expiresIn = DefaultVendorAuthCodeExpiresIn
}
return &VendorAuthCodeResult{AuthCode: authCode, ExpiresIn: expiresIn}, nil
}
+178
View File
@@ -0,0 +1,178 @@
// Copyright 2026 Alibaba Group
// Licensed under the Apache License, Version 2.0 (the "License");
// you may not use this file except in compliance with the License.
// You may obtain a copy of the License at
//
// http://www.apache.org/licenses/LICENSE-2.0
//
// Unless required by applicable law or agreed to in writing, software
// distributed under the License is distributed on an "AS IS" BASIS,
// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
// See the License for the specific language governing permissions and
// limitations under the License.
package auth
import (
"context"
"encoding/json"
"errors"
"io"
"net/http"
"net/http/httptest"
"strings"
"testing"
)
func TestFetchVendorAuthCodeSuccessEnvelope(t *testing.T) {
var gotBody map[string]string
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
if r.URL.Path != MCPVendorAuthCodePath {
t.Fatalf("path = %q, want %s", r.URL.Path, MCPVendorAuthCodePath)
}
if r.Method != http.MethodPost {
t.Fatalf("method = %s, want POST", r.Method)
}
if got := r.Header.Get("x-user-access-token"); got != "user-token" {
t.Fatalf("x-user-access-token = %q", got)
}
if got := r.Header.Get("x-dws-client-id"); got != "dws-client" {
t.Fatalf("x-dws-client-id = %q", got)
}
if got := r.Header.Get("x-dws-cli-version"); got != "1.2.3" {
t.Fatalf("x-dws-cli-version = %q", got)
}
if err := json.NewDecoder(r.Body).Decode(&gotBody); err != nil {
t.Fatalf("decode body: %v", err)
}
w.Header().Set("Cache-Control", "no-store")
_, _ = io.WriteString(w, `{"authCode":"tmp-code","expiresIn":120}`)
}))
defer srv.Close()
got, err := FetchVendorAuthCode(context.Background(), VendorAuthCodeInput{
AccessToken: "user-token",
ClientID: "dws-client",
CLIVersion: "1.2.3",
Vendor: "SafeChat",
CorpID: "dingxxxxxxxxxxxx",
HTTPClient: srv.Client(),
BaseURL: srv.URL,
})
if err != nil {
t.Fatalf("FetchVendorAuthCode() = %v", err)
}
if got.AuthCode != "tmp-code" || got.ExpiresIn != 120 {
t.Fatalf("result = %+v", got)
}
if gotBody["vendor"] != "safechat" || gotBody["corpId"] != "dingxxxxxxxxxxxx" {
t.Fatalf("posted body = %v", gotBody)
}
if _, ok := gotBody["redirectURI"]; ok {
t.Fatalf("posted redirectURI, body = %v", gotBody)
}
if _, ok := gotBody["domain"]; ok {
t.Fatalf("posted domain, body = %v", gotBody)
}
}
func TestFetchVendorAuthCodeParsesAlways200ServiceResult(t *testing.T) {
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) {
w.WriteHeader(http.StatusOK)
_, _ = io.WriteString(w, `{"success":false,"errorCode":"VENDOR_NOT_ENABLED","errorMsg":"not installed"}`)
}))
defer srv.Close()
_, err := FetchVendorAuthCode(context.Background(), VendorAuthCodeInput{
AccessToken: "user-token",
ClientID: "dws-client",
Vendor: "safechat",
CorpID: "dingxxxxxxxxxxxx",
HTTPClient: srv.Client(),
BaseURL: srv.URL,
})
var verr *VendorAuthCodeError
if !errors.As(err, &verr) || verr.Code != VendorAuthCodeVendorNotEnabled {
t.Fatalf("error = %v, want VENDOR_NOT_ENABLED", err)
}
if verr.Retryable() {
t.Fatal("VENDOR_NOT_ENABLED must not be retryable")
}
}
func TestFetchVendorAuthCodeRequiresLocalFields(t *testing.T) {
_, err := FetchVendorAuthCode(context.Background(), VendorAuthCodeInput{Vendor: "safechat", CorpID: "ding"})
var verr *VendorAuthCodeError
if !errors.As(err, &verr) || verr.Code != VendorAuthCodeParamError {
t.Fatalf("error = %v, want PARAM_ERROR", err)
}
}
func TestFetchVendorAuthCodeKeepsHTTPStatusOnNon200(t *testing.T) {
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) {
http.Error(w, "oops", http.StatusBadGateway)
}))
defer srv.Close()
_, err := FetchVendorAuthCode(context.Background(), VendorAuthCodeInput{
AccessToken: "user-token",
ClientID: "dws-client",
Vendor: "safechat",
CorpID: "dingxxxxxxxxxxxx",
HTTPClient: srv.Client(),
BaseURL: srv.URL,
})
var statusErr *HTTPStatusError
if !errors.As(err, &statusErr) || statusErr.StatusCode != http.StatusBadGateway {
t.Fatalf("error = %v, want HTTP 502", err)
}
}
func TestParseVendorAuthCodeResponseDefaultsExpiresIn(t *testing.T) {
got, err := parseVendorAuthCodeResponse([]byte(`{"authCode":"x"}`))
if err != nil {
t.Fatalf("parse: %v", err)
}
if got.ExpiresIn != DefaultVendorAuthCodeExpiresIn {
t.Fatalf("expiresIn = %d, want %d", got.ExpiresIn, DefaultVendorAuthCodeExpiresIn)
}
}
func TestVendorAuthCodeErrorRetryable(t *testing.T) {
for _, code := range []string{VendorAuthCodeTokenInvalid, VendorAuthCodeRateLimited, VendorAuthCodeInternalError} {
if !(&VendorAuthCodeError{Code: code}).Retryable() {
t.Fatalf("%s should be retryable", code)
}
}
if (&VendorAuthCodeError{Code: VendorAuthCodeOrgMismatch}).Retryable() {
t.Fatal("ORG_MISMATCH must not be retryable")
}
}
func TestFetchVendorAuthCodeDoesNotSendBlankCLIVersionHeader(t *testing.T) {
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
if got := r.Header.Get("x-dws-cli-version"); got != "" {
t.Fatalf("x-dws-cli-version = %q, want empty", got)
}
_, _ = io.WriteString(w, `{"authCode":"tmp-code","expiresIn":90}`)
}))
defer srv.Close()
got, err := FetchVendorAuthCode(context.Background(), VendorAuthCodeInput{
AccessToken: "user-token",
ClientID: "dws-client",
Vendor: "safechat",
CorpID: "dingxxxxxxxxxxxx",
HTTPClient: srv.Client(),
BaseURL: srv.URL,
})
if err != nil {
t.Fatalf("FetchVendorAuthCode() = %v", err)
}
if got.ExpiresIn != 90 {
t.Fatalf("expiresIn = %d, want 90 from response", got.ExpiresIn)
}
if strings.Contains(got.AuthCode, "redirect") {
t.Fatalf("unexpected code %q", got.AuthCode)
}
}
+40
View File
@@ -3,6 +3,7 @@
package keychain
import (
"bytes"
"encoding/base64"
"errors"
"os"
@@ -246,6 +247,45 @@ func TestCrossPlatformCoverageDarwinDEKKeyringEdges(t *testing.T) {
if got, err := getOrCreateDEK("generate-missing"); err != nil || len(got) != dekBytes {
t.Fatalf("generate missing = %d, %v", len(got), err)
}
existing := bytesOf(7, dekBytes)
setCalls := 0
keyringGet = func(string, string) (string, error) {
return base64.StdEncoding.EncodeToString(existing), nil
}
keyringSet = func(string, string, string) error {
setCalls++
return errors.New("duplicate write must not replace an existing DEK")
}
got, err := getOrCreateDEK("reuse-existing")
if err != nil || !bytes.Equal(got, existing) {
t.Fatalf("reuse existing = %d, %v", len(got), err)
}
if setCalls != 0 {
t.Fatalf("reuse existing set calls = %d, want 0", setCalls)
}
gets := 0
setCalls = 0
keyringGet = func(string, string) (string, error) {
gets++
if gets == 1 {
return "", keyring.ErrNotFound
}
return encoded, nil
}
keyringSet = func(string, string, string) error {
setCalls++
return errors.New("already exists")
}
got, err = getOrCreateDEK("create-race")
if err != nil || !bytes.Equal(got, valid) {
t.Fatalf("create race = %d, %v", len(got), err)
}
if setCalls != 1 {
t.Fatalf("create race set calls = %d, want 1", setCalls)
}
keyringGet = func(string, string) (string, error) { return "", keyring.ErrNotFound }
keychainRandRead = func([]byte) (int, error) { return 0, errKeychainInjected }
if _, err := getOrCreateDEK("rand"); err == nil {
t.Fatal("rand error expected")
+37 -26
View File
@@ -237,6 +237,26 @@ func getSystemDEKReadOnly(service string) ([]byte, error) {
return key, err
}
func decodeSystemDEK(encodedKey string) ([]byte, bool) {
key, err := base64.StdEncoding.DecodeString(encodedKey)
return key, err == nil && len(key) == dekBytes
}
func readStoredSystemDEK(service string, runtime darwinKeychainRuntime) ([]byte, error) {
encodedKey, err := runtime.get(service, "dek")
if err == nil {
key, ok := decodeSystemDEK(encodedKey)
if ok {
return key, nil
}
return nil, fmt.Errorf("read DEK from macOS Keychain: %w", ErrDEKMissing)
}
if errors.Is(err, keyring.ErrNotFound) {
return nil, fmt.Errorf("read DEK from macOS Keychain: %w", ErrDEKMissing)
}
return nil, NewUnavailableError("read DEK from macOS Keychain", err)
}
func getSystemDEKReadOnlyWithRuntime(service string, runtime darwinKeychainRuntime) ([]byte, error, <-chan struct{}) {
if err := runtime.checkAvailable(); err != nil {
return nil, err, finishedDarwinKeychainWorker()
@@ -244,18 +264,7 @@ func getSystemDEKReadOnlyWithRuntime(service string, runtime darwinKeychainRunti
const operation = "read DEK from macOS Keychain"
worker := startDarwinKeychainWorker(operation, func() ([]byte, error) {
encodedKey, err := runtime.get(service, "dek")
if err == nil {
key, decodeErr := base64.StdEncoding.DecodeString(encodedKey)
if decodeErr == nil && len(key) == dekBytes {
return key, nil
}
return nil, fmt.Errorf("read DEK from macOS Keychain: %w", ErrDEKMissing)
}
if errors.Is(err, keyring.ErrNotFound) {
return nil, fmt.Errorf("read DEK from macOS Keychain: %w", ErrDEKMissing)
}
return nil, NewUnavailableError("read DEK from macOS Keychain", err)
return readStoredSystemDEK(service, runtime)
})
return waitDarwinKeychainWorker(runtime.timeout, operation, worker)
@@ -276,28 +285,30 @@ func getOrCreateDEKWithRuntime(service string, runtime darwinKeychainRuntime) ([
const operation = "read or create DEK in macOS Keychain"
worker := startDarwinKeychainWorker(operation, func() ([]byte, error) {
// Try to get existing DEK from system Keychain
encodedKey, err := runtime.get(service, "dek")
key, err := readStoredSystemDEK(service, runtime)
if err == nil {
key, decodeErr := base64.StdEncoding.DecodeString(encodedKey)
if decodeErr == nil && len(key) == dekBytes {
return key, nil
}
} else if !errors.Is(err, keyring.ErrNotFound) {
return nil, NewUnavailableError("read DEK from macOS Keychain", err)
return key, nil
}
if !IsDEKMissing(err) {
return nil, err
}
// Generate new DEK if not found or invalid
key := make([]byte, dekBytes)
// Generate a candidate only when the slot is empty or unreadable.
// Concurrent writers must not replace a DEK another process just stored.
key = make([]byte, dekBytes)
if _, randErr := runtime.randRead(key); randErr != nil {
return nil, randErr
}
// Store in system Keychain
encodedKey = base64.StdEncoding.EncodeToString(key)
if setErr := runtime.set(service, "dek", encodedKey); setErr != nil {
if setErr := runtime.set(service, "dek", base64.StdEncoding.EncodeToString(key)); setErr != nil {
existing, getErr := readStoredSystemDEK(service, runtime)
if getErr == nil {
return existing, nil
}
return nil, NewUnavailableError("store DEK in macOS Keychain", setErr)
}
if existing, getErr := readStoredSystemDEK(service, runtime); getErr == nil {
return existing, nil
}
return key, nil
})
+127
View File
@@ -0,0 +1,127 @@
// Copyright 2026 Alibaba Group
// Licensed under the Apache License, Version 2.0 (the "License");
// you may not use this file except in compliance with the License.
// You may obtain a copy of the License at
//
// http://www.apache.org/licenses/LICENSE-2.0
//
// Unless required by applicable law or agreed to in writing, software
// distributed under the License is distributed on an "AS IS" BASIS,
// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
// See the License for the specific language governing permissions and
// limitations under the License.
package msgcrypto
import (
"context"
"errors"
"fmt"
"sync"
"time"
)
// DefaultAuthCodeTTL is the unconsumed-cache window for a freshly minted
// vendor authCode. The portal issues codes with expiresIn=120s and they are
// one-shot, so the default stays under that server window. Prefer not wrapping
// PortalAuthCode in CachedAuthCode: mint in goProxy and discard after the key
// request.
const DefaultAuthCodeTTL = 90 * time.Second
// ErrNoAuthCode means the provider returned an empty code without an error.
var ErrNoAuthCode = errors.New("msgcrypto: auth code provider returned an empty code")
// AuthCodeProvider yields a DingTalk 免登 authCode for key-server
// authentication. DWS does not mint the code itself, so integrations inject an
// implementation. The backend calls this only from the vendor goProxy
// callback, never on every encrypt or decrypt.
//
// Implementations must be safe for concurrent use; the backend may call this
// from a CGO callback while an encrypt or decrypt call is in flight.
type AuthCodeProvider interface {
AuthCode(ctx context.Context) (string, error)
}
// CorpAuthCodeProvider mints a code for a specific organization. PortalAuthCode
// implements this so goProxy can pass the C library's corpID. Domain and
// redirectURI are never part of this call.
type CorpAuthCodeProvider interface {
AuthCodeProvider
AuthCodeForCorp(ctx context.Context, corpID string) (string, error)
}
// AuthCodeFunc adapts a function to AuthCodeProvider.
type AuthCodeFunc func(ctx context.Context) (string, error)
// AuthCode calls f.
func (f AuthCodeFunc) AuthCode(ctx context.Context) (string, error) { return f(ctx) }
// StaticAuthCode returns a provider that always yields code. It is meant for
// tests and manual integration runs; a static code stops working once the
// server-side five-minute window closes.
func StaticAuthCode(code string) AuthCodeProvider {
return AuthCodeFunc(func(context.Context) (string, error) {
if code == "" {
return "", ErrNoAuthCode
}
return code, nil
})
}
// CachedAuthCode memoises an AuthCodeProvider for a TTL so a burst of key
// requests does not trigger one upstream call each.
type CachedAuthCode struct {
provider AuthCodeProvider
ttl time.Duration
now func() time.Time
mu sync.Mutex
code string
expiresAt time.Time
}
// NewCachedAuthCode wraps provider with a TTL cache. A ttl of zero or less
// selects DefaultAuthCodeTTL.
func NewCachedAuthCode(provider AuthCodeProvider, ttl time.Duration) *CachedAuthCode {
if ttl <= 0 {
ttl = DefaultAuthCodeTTL
}
return &CachedAuthCode{provider: provider, ttl: ttl, now: time.Now}
}
// AuthCode returns the cached code when it is still fresh, otherwise fetches a
// new one. A failed fetch leaves no stale value behind.
func (c *CachedAuthCode) AuthCode(ctx context.Context) (string, error) {
if c.provider == nil {
return "", ErrNoAuthCodeProvider
}
c.mu.Lock()
defer c.mu.Unlock()
if c.code != "" && c.now().Before(c.expiresAt) {
return c.code, nil
}
code, err := c.provider.AuthCode(ctx)
if err != nil {
c.code, c.expiresAt = "", time.Time{}
return "", fmt.Errorf("msgcrypto: fetch auth code: %w", err)
}
if code == "" {
c.code, c.expiresAt = "", time.Time{}
return "", ErrNoAuthCode
}
c.code = code
c.expiresAt = c.now().Add(c.ttl)
return code, nil
}
// Invalidate drops the cached code so the next AuthCode call refetches. The
// backend calls this after the key server rejects a code.
func (c *CachedAuthCode) Invalidate() {
c.mu.Lock()
c.code, c.expiresAt = "", time.Time{}
c.mu.Unlock()
}
+228
View File
@@ -0,0 +1,228 @@
// Copyright 2026 Alibaba Group
// Licensed under the Apache License, Version 2.0 (the "License");
// you may not use this file except in compliance with the License.
// You may obtain a copy of the License at
//
// http://www.apache.org/licenses/LICENSE-2.0
//
// Unless required by applicable law or agreed to in writing, software
// distributed under the License is distributed on an "AS IS" BASIS,
// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
// See the License for the specific language governing permissions and
// limitations under the License.
package msgcrypto
import (
"context"
"errors"
"sync"
"testing"
"time"
)
// countingProvider hands out a fresh code per call and records how often it was
// asked, so cache behaviour can be asserted.
type countingProvider struct {
mu sync.Mutex
calls int
code string
err error
}
func (p *countingProvider) AuthCode(context.Context) (string, error) {
p.mu.Lock()
defer p.mu.Unlock()
p.calls++
if p.err != nil {
return "", p.err
}
if p.code != "" {
return p.code, nil
}
return "code-" + string(rune('a'+p.calls-1)), nil
}
// callCount reports the number of upstream fetches.
func (p *countingProvider) callCount() int {
p.mu.Lock()
defer p.mu.Unlock()
return p.calls
}
func TestAuthCodeFuncAdaptsFunction(t *testing.T) {
provider := AuthCodeFunc(func(context.Context) (string, error) { return "abc", nil })
code, err := provider.AuthCode(context.Background())
if err != nil || code != "abc" {
t.Fatalf("AuthCode() = %q, %v; want abc, nil", code, err)
}
}
func TestStaticAuthCodeReturnsCode(t *testing.T) {
code, err := StaticAuthCode("fixed").AuthCode(context.Background())
if err != nil || code != "fixed" {
t.Fatalf("AuthCode() = %q, %v; want fixed, nil", code, err)
}
}
func TestStaticAuthCodeRejectsEmptyCode(t *testing.T) {
_, err := StaticAuthCode("").AuthCode(context.Background())
if !errors.Is(err, ErrNoAuthCode) {
t.Fatalf("AuthCode() = %v, want ErrNoAuthCode", err)
}
}
func TestCachedAuthCodeReusesCodeWithinTTL(t *testing.T) {
provider := &countingProvider{code: "same"}
cache := NewCachedAuthCode(provider, time.Minute)
for i := 0; i < 5; i++ {
code, err := cache.AuthCode(context.Background())
if err != nil {
t.Fatalf("AuthCode() #%d = %v", i+1, err)
}
if code != "same" {
t.Fatalf("AuthCode() #%d = %q, want same", i+1, code)
}
}
if got := provider.callCount(); got != 1 {
t.Fatalf("upstream called %d times, want 1 (the code must be cached)", got)
}
}
func TestCachedAuthCodeRefetchesAfterTTL(t *testing.T) {
provider := &countingProvider{}
cache := NewCachedAuthCode(provider, time.Minute)
now := time.Now()
cache.now = func() time.Time { return now }
first, err := cache.AuthCode(context.Background())
if err != nil {
t.Fatalf("first AuthCode() = %v", err)
}
// Move past the TTL. The DingTalk code expires server-side, so a stale
// one must not be reused.
now = now.Add(time.Minute + time.Second)
second, err := cache.AuthCode(context.Background())
if err != nil {
t.Fatalf("second AuthCode() = %v", err)
}
if first == second {
t.Fatalf("AuthCode() returned the same code %q after the TTL expired", first)
}
if got := provider.callCount(); got != 2 {
t.Fatalf("upstream called %d times, want 2", got)
}
}
func TestCachedAuthCodeDefaultTTLIsUnderServerWindow(t *testing.T) {
// Portal vendorAuthCode expiresIn is 120s and the code is one-shot.
// The unconsumed-cache window must stay under that server lifetime.
if DefaultAuthCodeTTL >= 120*time.Second {
t.Fatalf("DefaultAuthCodeTTL = %v, want less than the 120s portal expiresIn", DefaultAuthCodeTTL)
}
cache := NewCachedAuthCode(&countingProvider{}, 0)
if cache.ttl != DefaultAuthCodeTTL {
t.Fatalf("ttl = %v, want DefaultAuthCodeTTL %v", cache.ttl, DefaultAuthCodeTTL)
}
}
func TestCachedAuthCodeNegativeTTLFallsBackToDefault(t *testing.T) {
cache := NewCachedAuthCode(&countingProvider{}, -time.Second)
if cache.ttl != DefaultAuthCodeTTL {
t.Fatalf("ttl = %v, want DefaultAuthCodeTTL %v", cache.ttl, DefaultAuthCodeTTL)
}
}
func TestCachedAuthCodePropagatesUpstreamError(t *testing.T) {
wantErr := errors.New("token service down")
cache := NewCachedAuthCode(&countingProvider{err: wantErr}, time.Minute)
_, err := cache.AuthCode(context.Background())
if !errors.Is(err, wantErr) {
t.Fatalf("AuthCode() = %v, want it to wrap %v", err, wantErr)
}
}
func TestCachedAuthCodeDoesNotCacheFailures(t *testing.T) {
provider := &countingProvider{err: errors.New("transient")}
cache := NewCachedAuthCode(provider, time.Minute)
if _, err := cache.AuthCode(context.Background()); err == nil {
t.Fatal("AuthCode() = nil error, want failure")
}
provider.mu.Lock()
provider.err = nil
provider.code = "recovered"
provider.mu.Unlock()
code, err := cache.AuthCode(context.Background())
if err != nil {
t.Fatalf("AuthCode() after recovery = %v", err)
}
if code != "recovered" {
t.Fatalf("AuthCode() = %q, want recovered (a failure must not be cached)", code)
}
}
func TestCachedAuthCodeRejectsEmptyUpstreamCode(t *testing.T) {
// A provider that reports success with no code is a bug upstream; the
// cache must surface it instead of caching an unusable value.
cache := NewCachedAuthCode(AuthCodeFunc(func(context.Context) (string, error) {
return "", nil
}), time.Minute)
if _, err := cache.AuthCode(context.Background()); !errors.Is(err, ErrNoAuthCode) {
t.Fatalf("AuthCode() = %v, want ErrNoAuthCode", err)
}
}
func TestCachedAuthCodeInvalidateForcesRefetch(t *testing.T) {
provider := &countingProvider{}
cache := NewCachedAuthCode(provider, time.Hour)
if _, err := cache.AuthCode(context.Background()); err != nil {
t.Fatalf("first AuthCode() = %v", err)
}
cache.Invalidate()
if _, err := cache.AuthCode(context.Background()); err != nil {
t.Fatalf("second AuthCode() = %v", err)
}
if got := provider.callCount(); got != 2 {
t.Fatalf("upstream called %d times, want 2 after Invalidate", got)
}
}
func TestCachedAuthCodeWithoutProviderReportsMissingProvider(t *testing.T) {
cache := NewCachedAuthCode(nil, time.Minute)
if _, err := cache.AuthCode(context.Background()); !errors.Is(err, ErrNoAuthCodeProvider) {
t.Fatalf("AuthCode() = %v, want ErrNoAuthCodeProvider", err)
}
}
func TestCachedAuthCodeIsSafeForConcurrentUse(t *testing.T) {
// The backend may ask for a code from a CGO callback while another
// operation is in flight, so concurrent access must not race.
provider := &countingProvider{code: "shared"}
cache := NewCachedAuthCode(provider, time.Hour)
var wg sync.WaitGroup
for i := 0; i < 32; i++ {
wg.Add(1)
go func() {
defer wg.Done()
if code, err := cache.AuthCode(context.Background()); err != nil || code != "shared" {
t.Errorf("AuthCode() = %q, %v; want shared, nil", code, err)
}
}()
}
wg.Wait()
if got := provider.callCount(); got != 1 {
t.Fatalf("upstream called %d times, want 1", got)
}
}
+159
View File
@@ -0,0 +1,159 @@
// Copyright 2026 Alibaba Group
// Licensed under the Apache License, Version 2.0 (the "License");
// you may not use this file except in compliance with the License.
// You may obtain a copy of the License at
//
// http://www.apache.org/licenses/LICENSE-2.0
//
// Unless required by applicable law or agreed to in writing, software
// distributed under the License is distributed on an "AS IS" BASIS,
// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
// See the License for the specific language governing permissions and
// limitations under the License.
// The constraint below must stay in sync with cipher_stub.go, which negates it
// verbatim. It encodes the platforms the vendor ships a libsafechat.a for:
// darwin and linux on amd64/arm64, plus windows/amd64. windows/arm64 is
// deliberately excluded because the vendor has not delivered that static
// library, and DWS does release that target.
//go:build safechat && cgo && (((darwin || linux) && (amd64 || arm64)) || (windows && amd64))
package msgcrypto
import (
"context"
"errors"
"fmt"
"sync"
safechat "safechat-go-sdk"
)
// BackendVersion identifies the compiled-in vendor SDK.
const BackendVersion = "safechat " + safechat.Version
// Available reports that this binary carries the SafeChat backend.
func Available() bool { return true }
// safechatCipher adapts the vendor client to Cipher.
//
// The vendor client serialises its own C calls internally, so this type adds no
// further locking. Auth codes are minted only from AuthCodeHook, which the
// vendor SDK calls inside goProxy when a key is actually missing.
type safechatCipher struct {
client *safechat.Client
codes AuthCodeProvider
allowedHost string
mu sync.Mutex
lastCodeErr error
}
// newBackend starts the vendor client against cfg's keystore.
//
// A warm keystore serves encrypt and decrypt without a key request, so no
// authCode is fetched at open time. The hook runs only if goProxy fires.
func newBackend(ctx context.Context, cfg Config) (Cipher, error) {
if err := ctx.Err(); err != nil {
return nil, err
}
var logf func(string, ...any)
if cfg.Debug {
logf = cfg.Logf
}
c := &safechatCipher{codes: cfg.AuthCode, allowedHost: cfg.AllowedRedirectHost}
client, err := safechat.New(safechat.Config{
DataPath: cfg.KeystoreDir,
UserID: cfg.UserID,
KeyServer: cfg.KeyServer,
MaxRetry: cfg.MaxRetry,
HTTPTimeout: cfg.HTTPTimeout,
Logger: newRedactingLogger(logf),
AuthCodeHook: c.authCodeHook,
})
if err != nil {
if errors.Is(err, safechat.ErrAlreadyInitialized) {
return nil, ErrAlreadyOpen
}
return nil, fmt.Errorf("msgcrypto: start safechat backend: %w", err)
}
c.client = client
return c, nil
}
// EncryptMessage encrypts plaintext and returns the vendor ciphertext.
func (c *safechatCipher) EncryptMessage(ctx context.Context, corpID, staffID string, plaintext []byte) ([]byte, error) {
if err := ctx.Err(); err != nil {
return nil, err
}
c.setLastCodeErr(nil)
out, err := c.client.EncryptMsg(corpID, staffID, plaintext)
if err != nil {
return nil, c.explain("encrypt", corpID, err)
}
return out, nil
}
// DecryptMessage decrypts a vendor ciphertext.
func (c *safechatCipher) DecryptMessage(ctx context.Context, corpID, staffID string, ciphertext []byte) ([]byte, error) {
if err := ctx.Err(); err != nil {
return nil, err
}
c.setLastCodeErr(nil)
out, err := c.client.DecryptMsg(corpID, staffID, ciphertext)
if err != nil {
return nil, c.explain("decrypt", corpID, err)
}
return out, nil
}
// Close releases the vendor client.
func (c *safechatCipher) Close() error {
c.client.Close()
return nil
}
// authCodeHook is invoked from goProxy immediately before the key request.
// domain is compared locally and never forwarded to portal. The returned
// code is used once by the SDK and is not stored on the client.
func (c *safechatCipher) authCodeHook(corpID, domain string) (string, error) {
code, err := mintAuthCodeForProxy(c.codes, c.allowedHost, corpID, domain)
c.setLastCodeErr(err)
return code, err
}
func (c *safechatCipher) setLastCodeErr(err error) {
c.mu.Lock()
c.lastCodeErr = err
c.mu.Unlock()
}
func (c *safechatCipher) lastAuthCodeErr() error {
c.mu.Lock()
defer c.mu.Unlock()
return c.lastCodeErr
}
// explain turns a vendor error into an actionable one, folding in a failed
// goProxy authCode mint and the admin-restricted case.
func (c *safechatCipher) explain(op, corpID string, opErr error) error {
if c.client.IsBlocked(corpID) {
return fmt.Errorf("msgcrypto: %s blocked: the organization's key is restricted by its administrator: %w", op, opErr)
}
// A key fetch was needed but we had no usable code: that is the real
// cause, so report both.
if codeErr := c.lastAuthCodeErr(); codeErr != nil {
return fmt.Errorf("msgcrypto: %s failed and no usable auth code was available: %w (auth code error: %v)", op, opErr, codeErr)
}
if errors.Is(opErr, safechat.ErrMaxRetryExceeded) {
invalidateAuthCode(c.codes)
return fmt.Errorf("msgcrypto: %s failed: key material never became available: %w", op, opErr)
}
return fmt.Errorf("msgcrypto: %s failed: %w", op, opErr)
}
+161
View File
@@ -0,0 +1,161 @@
// Copyright 2026 Alibaba Group
// Licensed under the Apache License, Version 2.0 (the "License");
// you may not use this file except in compliance with the License.
// You may obtain a copy of the License at
//
// http://www.apache.org/licenses/LICENSE-2.0
//
// Unless required by applicable law or agreed to in writing, software
// distributed under the License is distributed on an "AS IS" BASIS,
// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
// See the License for the specific language governing permissions and
// limitations under the License.
// Keep this constraint in sync with cipher_safechat.go.
//go:build safechat && cgo && (((darwin || linux) && (amd64 || arm64)) || (windows && amd64))
package msgcrypto
import (
"context"
"errors"
"os"
"path/filepath"
"strings"
"testing"
)
// These tests exercise the real vendor backend, which means they link
// libsafechat.a and initialise the C library. They never reach the key server:
// an authCode is only spent when a key is actually fetched, and asserting that
// encryption fails without one is exactly the behaviour we want pinned.
func TestBackendIsReportedAvailable(t *testing.T) {
if !Available() {
t.Fatal("Available() = false in a safechat build")
}
if BackendVersion == "" {
t.Fatal("BackendVersion is empty in a safechat build")
}
if !strings.Contains(BackendVersion, "safechat") {
t.Fatalf("BackendVersion = %q, want it to name the vendor SDK", BackendVersion)
}
}
func TestOpenInitialisesCLibraryAndClosesCleanly(t *testing.T) {
dir := filepath.Join(t.TempDir(), "keystore")
cipher, err := Open(context.Background(), Config{
KeystoreDir: dir,
AuthCode: StaticAuthCode("placeholder-code"),
KeyServer: "https://key.example.test",
})
if err != nil {
t.Fatalf("Open() = %v, want the C library to initialise", err)
}
if info, statErr := os.Stat(dir); statErr != nil || !info.IsDir() {
t.Fatalf("Open did not prepare the keystore dir: %v", statErr)
}
if err := cipher.Close(); err != nil {
t.Fatalf("Close() = %v", err)
}
// The slot must be free again, otherwise a second Open in the same
// process would be refused forever.
second, err := Open(context.Background(), Config{
KeystoreDir: dir,
AuthCode: StaticAuthCode("placeholder-code"),
KeyServer: "https://key.example.test",
})
if err != nil {
t.Fatalf("second Open() after Close = %v, want success", err)
}
if err := second.Close(); err != nil {
t.Fatalf("second Close() = %v", err)
}
}
func TestOpenRefusesConcurrentSecondCipher(t *testing.T) {
dir := filepath.Join(t.TempDir(), "keystore")
first, err := Open(context.Background(), Config{
KeystoreDir: dir,
AuthCode: StaticAuthCode("placeholder-code"),
KeyServer: "https://key.example.test",
})
if err != nil {
t.Fatalf("Open() = %v", err)
}
defer first.Close()
_, err = Open(context.Background(), Config{
KeystoreDir: filepath.Join(t.TempDir(), "keystore"),
AuthCode: StaticAuthCode("placeholder-code"),
KeyServer: "https://key.example.test",
})
if !errors.Is(err, ErrAlreadyOpen) {
t.Fatalf("second Open() = %v, want ErrAlreadyOpen (the C library keeps global state)", err)
}
}
func TestOpenHonoursCancelledContext(t *testing.T) {
ctx, cancel := context.WithCancel(context.Background())
cancel()
_, err := Open(ctx, Config{
KeystoreDir: filepath.Join(t.TempDir(), "keystore"),
AuthCode: StaticAuthCode("placeholder-code"),
KeyServer: "https://key.example.test",
})
if !errors.Is(err, context.Canceled) {
t.Fatalf("Open() with a cancelled context = %v, want context.Canceled", err)
}
}
func TestEncryptWithoutUsableKeyReportsAuthCodeCause(t *testing.T) {
// A cold keystore forces a key request. With no reachable key server the
// operation must fail with a message that names the auth code, rather
// than a bare vendor return code.
cipher, err := Open(context.Background(), Config{
KeystoreDir: filepath.Join(t.TempDir(), "keystore"),
AuthCode: AuthCodeFunc(func(context.Context) (string, error) {
return "", errors.New("no code available in test")
}),
KeyServer: "https://key.example.test",
})
if err != nil {
t.Fatalf("Open() = %v", err)
}
defer cipher.Close()
_, err = cipher.EncryptMessage(context.Background(), "test-corp", "test-staff", []byte("hello"))
if err == nil {
t.Skip("the environment served a key without an auth code; nothing to assert")
}
if !strings.Contains(err.Error(), "auth code") {
t.Fatalf("EncryptMessage() = %v, want the error to name the auth code cause", err)
}
}
func TestCipherRejectsBadArgumentsBeforeCallingC(t *testing.T) {
cipher, err := Open(context.Background(), Config{
KeystoreDir: filepath.Join(t.TempDir(), "keystore"),
AuthCode: StaticAuthCode("placeholder-code"),
KeyServer: "https://key.example.test",
})
if err != nil {
t.Fatalf("Open() = %v", err)
}
defer cipher.Close()
if _, err := cipher.EncryptMessage(context.Background(), "", "staff", []byte("x")); !errors.Is(err, ErrNoCorpID) {
t.Fatalf("EncryptMessage() with no corpID = %v, want ErrNoCorpID", err)
}
// The vendor SDK dereferences the first byte of the payload, so an empty
// slice must never reach it.
if _, err := cipher.DecryptMessage(context.Background(), "corp", "staff", nil); !errors.Is(err, ErrEmptyPayload) {
t.Fatalf("DecryptMessage() with no payload = %v, want ErrEmptyPayload", err)
}
}
+34
View File
@@ -0,0 +1,34 @@
// Copyright 2026 Alibaba Group
// Licensed under the Apache License, Version 2.0 (the "License");
// you may not use this file except in compliance with the License.
// You may obtain a copy of the License at
//
// http://www.apache.org/licenses/LICENSE-2.0
//
// Unless required by applicable law or agreed to in writing, software
// distributed under the License is distributed on an "AS IS" BASIS,
// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
// See the License for the specific language governing permissions and
// limitations under the License.
// The constraint below is the exact negation of the one in cipher_safechat.go;
// change both together. This file covers every default DWS build: the release
// binaries are cross-compiled with CGO_ENABLED=0, and windows/arm64 has no
// vendor static library even when the tag is set.
//go:build !(safechat && cgo && (((darwin || linux) && (amd64 || arm64)) || (windows && amd64)))
package msgcrypto
import "context"
// BackendVersion is empty because no backend is compiled in.
const BackendVersion = ""
// Available reports that this binary has no SafeChat backend, so callers should
// not offer message encryption.
func Available() bool { return false }
// newBackend always fails here. Open checks Available first, so this exists to
// keep the package compiling and to fail safe if that check is ever bypassed.
func newBackend(context.Context, Config) (Cipher, error) { return nil, ErrUnavailable }
+54
View File
@@ -0,0 +1,54 @@
// Copyright 2026 Alibaba Group
// Licensed under the Apache License, Version 2.0 (the "License");
// you may not use this file except in compliance with the License.
// You may obtain a copy of the License at
//
// http://www.apache.org/licenses/LICENSE-2.0
//
// Unless required by applicable law or agreed to in writing, software
// distributed under the License is distributed on an "AS IS" BASIS,
// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
// See the License for the specific language governing permissions and
// limitations under the License.
package msgcrypto
import "context"
// mintAuthCodeForProxy is the goProxy-only mint path. domain is compared
// locally and never forwarded. A successful mint invalidates any unconsumed
// cache so a one-shot code cannot be reused.
func mintAuthCodeForProxy(codes AuthCodeProvider, allowedHost, corpID, domain string) (string, error) {
if err := matchRedirectHost(domain, allowedHost); err != nil {
return "", err
}
if codes == nil {
return "", ErrNoAuthCodeProvider
}
var (
code string
err error
)
if provider, ok := codes.(CorpAuthCodeProvider); ok {
code, err = provider.AuthCodeForCorp(context.Background(), corpID)
} else {
code, err = codes.AuthCode(context.Background())
}
if err != nil {
invalidateAuthCode(codes)
return "", err
}
if code == "" {
invalidateAuthCode(codes)
return "", ErrNoAuthCode
}
invalidateAuthCode(codes)
return code, nil
}
func invalidateAuthCode(codes AuthCodeProvider) {
if invalidator, ok := codes.(interface{ Invalidate() }); ok {
invalidator.Invalidate()
}
}
@@ -0,0 +1,99 @@
// Copyright 2026 Alibaba Group
// Licensed under the Apache License, Version 2.0 (the "License");
// you may not use this file except in compliance with the License.
// You may obtain a copy of the License at
//
// http://www.apache.org/licenses/LICENSE-2.0
//
// Unless required by applicable law or agreed to in writing, software
// distributed under the License is distributed on an "AS IS" BASIS,
// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
// See the License for the specific language governing permissions and
// limitations under the License.
package msgcrypto
import (
"context"
"errors"
"sync"
"testing"
"time"
)
type corpCountingProvider struct {
mu sync.Mutex
calls []string
code string
err error
}
func (p *corpCountingProvider) AuthCode(context.Context) (string, error) {
return p.AuthCodeForCorp(context.Background(), "")
}
func (p *corpCountingProvider) AuthCodeForCorp(_ context.Context, corpID string) (string, error) {
p.mu.Lock()
defer p.mu.Unlock()
p.calls = append(p.calls, corpID)
if p.err != nil {
return "", p.err
}
if p.code != "" {
return p.code, nil
}
return "code-for-" + corpID, nil
}
func TestMintAuthCodeForProxyUsesCorpProvider(t *testing.T) {
inner := &corpCountingProvider{}
code, err := mintAuthCodeForProxy(inner, "sso.anhei.test", "ding_corp", "https://sso.anhei.test/login")
if err != nil || code != "code-for-ding_corp" {
t.Fatalf("mint = %q, %v; want code-for-ding_corp, nil", code, err)
}
if len(inner.calls) != 1 || inner.calls[0] != "ding_corp" {
t.Fatalf("corpIDs = %v, want [ding_corp]", inner.calls)
}
}
func TestMintAuthCodeForProxyInvalidatesUnconsumedCache(t *testing.T) {
inner := &countingProvider{code: "once"}
cache := NewCachedAuthCode(inner, time.Hour)
if _, err := cache.AuthCode(context.Background()); err != nil {
t.Fatalf("seed cache: %v", err)
}
if got := inner.callCount(); got != 1 {
t.Fatalf("seed fetches = %d, want 1", got)
}
code, err := mintAuthCodeForProxy(cache, "sso.anhei.test", "ding_corp", "https://sso.anhei.test/login")
if err != nil || code != "once" {
t.Fatalf("mint = %q, %v; want once, nil", code, err)
}
// Cache was invalidated after spend; next mint hits upstream again.
if _, err := mintAuthCodeForProxy(cache, "sso.anhei.test", "ding_corp", "sso.anhei.test"); err != nil {
t.Fatalf("second mint: %v", err)
}
if got := inner.callCount(); got != 2 {
t.Fatalf("upstream calls = %d, want 2 (seed reused once, then refetch)", got)
}
}
func TestMintAuthCodeForProxyRejectsDomainMismatchWithoutFetching(t *testing.T) {
inner := &corpCountingProvider{code: "once"}
_, err := mintAuthCodeForProxy(inner, "sso.anhei.test", "ding_corp", "evil.example.test")
if !errors.Is(err, ErrRedirectHostMismatch) {
t.Fatalf("mint = %v, want ErrRedirectHostMismatch", err)
}
if got := len(inner.calls); got != 0 {
t.Fatalf("upstream called %d times on domain mismatch, want 0", got)
}
}
func TestMintAuthCodeForProxyRequiresProvider(t *testing.T) {
if _, err := mintAuthCodeForProxy(nil, "", "ding_corp", ""); !errors.Is(err, ErrNoAuthCodeProvider) {
t.Fatalf("mint = %v, want ErrNoAuthCodeProvider", err)
}
}
+83
View File
@@ -0,0 +1,83 @@
// Copyright 2026 Alibaba Group
// Licensed under the Apache License, Version 2.0 (the "License");
// you may not use this file except in compliance with the License.
// You may obtain a copy of the License at
//
// http://www.apache.org/licenses/LICENSE-2.0
//
// Unless required by applicable law or agreed to in writing, software
// distributed under the License is distributed on an "AS IS" BASIS,
// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
// See the License for the specific language governing permissions and
// limitations under the License.
package msgcrypto
import (
"fmt"
"net"
"net/url"
"strings"
)
// validateKeyServer requires an HTTPS URL with a host so the vendor C
// library cannot pick the key-request destination.
func validateKeyServer(raw string) error {
raw = strings.TrimSpace(raw)
if raw == "" {
return ErrNoKeyServer
}
u, err := url.Parse(raw)
if err != nil || u.Host == "" || u.Scheme == "" {
return fmt.Errorf("%w: %q", ErrInvalidKeyServer, raw)
}
if !strings.EqualFold(u.Scheme, "https") {
return fmt.Errorf("%w: %q", ErrKeyServerNotHTTPS, raw)
}
if hostnameOf(raw) == "" {
return fmt.Errorf("%w: %q", ErrInvalidKeyServer, raw)
}
return nil
}
// matchRedirectHost compares the goProxy domain to AllowedRedirectHost.
// Both sides are reduced to a hostname. An empty domain or an empty
// allowed host skips the check; the domain is never sent to portal.
func matchRedirectHost(domain, allowed string) error {
domain = strings.TrimSpace(domain)
allowed = strings.TrimSpace(allowed)
if domain == "" || allowed == "" {
return nil
}
got := hostnameOf(domain)
want := hostnameOf(allowed)
if got == "" || want == "" || got != want {
return fmt.Errorf("%w: got %q, want %q", ErrRedirectHostMismatch, got, want)
}
return nil
}
// hostnameOf returns the lower-cased hostname of a URL, host:port, or bare
// host. Path, query, userinfo and port are ignored.
func hostnameOf(raw string) string {
raw = strings.TrimSpace(raw)
if raw == "" {
return ""
}
if strings.Contains(raw, "://") {
u, err := url.Parse(raw)
if err == nil {
if host := strings.ToLower(u.Hostname()); host != "" {
return host
}
}
}
candidate := raw
if i := strings.IndexAny(candidate, "/?"); i >= 0 {
candidate = candidate[:i]
}
if host, _, err := net.SplitHostPort(candidate); err == nil {
return strings.ToLower(host)
}
return strings.ToLower(candidate)
}
+77
View File
@@ -0,0 +1,77 @@
// Copyright 2026 Alibaba Group
// Licensed under the Apache License, Version 2.0 (the "License");
// you may not use this file except in compliance with the License.
// You may obtain a copy of the License at
//
// http://www.apache.org/licenses/LICENSE-2.0
//
// Unless required by applicable law or agreed to in writing, software
// distributed under the License is distributed on an "AS IS" BASIS,
// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
// See the License for the specific language governing permissions and
// limitations under the License.
package msgcrypto
import (
"errors"
"testing"
)
func TestValidateKeyServerAcceptsHTTPS(t *testing.T) {
if err := validateKeyServer("https://key.example.test/v1"); err != nil {
t.Fatalf("validateKeyServer() = %v, want nil", err)
}
}
func TestValidateKeyServerRejectsEmpty(t *testing.T) {
if err := validateKeyServer(" "); !errors.Is(err, ErrNoKeyServer) {
t.Fatalf("validateKeyServer() = %v, want ErrNoKeyServer", err)
}
}
func TestValidateKeyServerRejectsHTTP(t *testing.T) {
if err := validateKeyServer("http://key.example.test"); !errors.Is(err, ErrKeyServerNotHTTPS) {
t.Fatalf("validateKeyServer() = %v, want ErrKeyServerNotHTTPS", err)
}
}
func TestValidateKeyServerRejectsBareHost(t *testing.T) {
if err := validateKeyServer("key.example.test"); !errors.Is(err, ErrInvalidKeyServer) {
t.Fatalf("validateKeyServer() = %v, want ErrInvalidKeyServer", err)
}
}
func TestMatchRedirectHostComparesHostOnly(t *testing.T) {
if err := matchRedirectHost("https://sso.anhei.test:443/login", "https://sso.anhei.test/path"); err != nil {
t.Fatalf("matchRedirectHost() = %v, want nil", err)
}
if err := matchRedirectHost("sso.anhei.test", "https://sso.anhei.test"); err != nil {
t.Fatalf("bare host match = %v, want nil", err)
}
}
func TestMatchRedirectHostSkipsWhenEitherSideEmpty(t *testing.T) {
if err := matchRedirectHost("", "https://sso.anhei.test"); err != nil {
t.Fatalf("empty domain = %v, want nil", err)
}
if err := matchRedirectHost("sso.anhei.test", ""); err != nil {
t.Fatalf("empty allowed host = %v, want nil", err)
}
}
func TestMatchRedirectHostRejectsMismatch(t *testing.T) {
err := matchRedirectHost("evil.example.test", "https://sso.anhei.test")
if !errors.Is(err, ErrRedirectHostMismatch) {
t.Fatalf("matchRedirectHost() = %v, want ErrRedirectHostMismatch", err)
}
}
func TestHostnameOfStripsPortAndPath(t *testing.T) {
if got := hostnameOf("https://SSO.Example.TEST:8443/login?x=1"); got != "sso.example.test" {
t.Fatalf("hostnameOf(url) = %q", got)
}
if got := hostnameOf("SSO.Example.TEST:8443/login"); got != "sso.example.test" {
t.Fatalf("hostnameOf(hostport) = %q", got)
}
}
+96
View File
@@ -0,0 +1,96 @@
// Copyright 2026 Alibaba Group
// Licensed under the Apache License, Version 2.0 (the "License");
// you may not use this file except in compliance with the License.
// You may obtain a copy of the License at
//
// http://www.apache.org/licenses/LICENSE-2.0
//
// Unless required by applicable law or agreed to in writing, software
// distributed under the License is distributed on an "AS IS" BASIS,
// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
// See the License for the specific language governing permissions and
// limitations under the License.
package msgcrypto
import "fmt"
// redactingLogger satisfies the vendor SDK's logger interface while making it
// impossible for the SDK to leak secrets into DWS output.
//
// This matters because the vendor SDK logs, at debug level, both the authCode
// it sends to the key server and the raw key-server response body, which
// carries key material. Rather than trying to enumerate and pattern-match
// every sensitive field, we drop every string argument and keep only its
// length. Numeric and boolean arguments pass through, so operators still get
// the useful diagnostics: HTTP status, payload sizes, timings and return
// codes. The format string is preserved so it stays clear which field was
// elided.
type redactingLogger struct {
logf func(format string, args ...any)
}
// newRedactingLogger returns a logger that forwards to logf, or nil when logf
// is nil so the SDK skips logging entirely.
func newRedactingLogger(logf func(format string, args ...any)) *redactingLogger {
if logf == nil {
return nil
}
return &redactingLogger{logf: logf}
}
// Debug forwards a redacted debug line.
func (l *redactingLogger) Debug(msg string, args ...interface{}) { l.emit("debug", msg, args) }
// Info forwards a redacted info line.
func (l *redactingLogger) Info(msg string, args ...interface{}) { l.emit("info", msg, args) }
// Error forwards a redacted error line.
func (l *redactingLogger) Error(msg string, args ...interface{}) { l.emit("error", msg, args) }
// emit rewrites args so no string value survives, then forwards the line.
func (l *redactingLogger) emit(level, msg string, args []interface{}) {
if l == nil || l.logf == nil {
return
}
l.logf("safechat[%s] "+msg, append([]any{level}, redactArgs(args)...)...)
}
// redactArgs replaces every string-like argument with a length marker and
// leaves other kinds intact.
func redactArgs(args []interface{}) []any {
out := make([]any, 0, len(args))
for _, arg := range args {
out = append(out, redactArg(arg))
}
return out
}
// redactArg elides a single argument's contents when it could carry a secret.
// Strings and byte slices are reduced to their length; errors are reduced to
// their type so a wrapped body preview cannot slip through; everything else
// (numbers, booleans, durations) is kept because it cannot carry key material.
func redactArg(arg any) any {
switch v := arg.(type) {
case string:
return redactedValue(len(v))
case []byte:
return redactedValue(len(v))
case fmt.Stringer:
// time.Duration and friends are Stringers, but so are opaque types
// that may embed a payload. Keep durations, elide the rest.
if _, ok := arg.(interface{ Nanoseconds() int64 }); ok {
return v.String()
}
return redactedValue(len(v.String()))
case error:
return fmt.Sprintf("<%T redacted>", v)
default:
return v
}
}
// redactedValue renders the placeholder used in place of elided content.
func redactedValue(n int) string {
return fmt.Sprintf("<redacted len=%d>", n)
}
+157
View File
@@ -0,0 +1,157 @@
// Copyright 2026 Alibaba Group
// Licensed under the Apache License, Version 2.0 (the "License");
// you may not use this file except in compliance with the License.
// You may obtain a copy of the License at
//
// http://www.apache.org/licenses/LICENSE-2.0
//
// Unless required by applicable law or agreed to in writing, software
// distributed under the License is distributed on an "AS IS" BASIS,
// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
// See the License for the specific language governing permissions and
// limitations under the License.
package msgcrypto
import (
"errors"
"fmt"
"strings"
"testing"
"time"
)
// captureLogf collects emitted lines for assertions.
func captureLogf(lines *[]string) func(string, ...any) {
return func(format string, args ...any) {
*lines = append(*lines, fmt.Sprintf(format, args...))
}
}
func TestNewRedactingLoggerReturnsNilWhenSinkIsNil(t *testing.T) {
// A nil logger makes the vendor SDK skip logging entirely, which is the
// safe default because it logs the authCode at debug level.
if got := newRedactingLogger(nil); got != nil {
t.Fatalf("newRedactingLogger(nil) = %v, want nil", got)
}
}
func TestRedactingLoggerHidesAuthCode(t *testing.T) {
var lines []string
logger := newRedactingLogger(captureLogf(&lines))
// This mirrors the vendor SDK's own debug line, which prints the code.
const secret = "abc123authcode"
logger.Debug("Code (auth_token, length=%d): %s", len(secret), secret)
if len(lines) != 1 {
t.Fatalf("got %d lines, want 1", len(lines))
}
if strings.Contains(lines[0], secret) {
t.Fatalf("log line leaked the auth code: %q", lines[0])
}
if !strings.Contains(lines[0], "redacted") {
t.Fatalf("log line = %q, want a redaction marker", lines[0])
}
}
func TestRedactingLoggerHidesKeyServerResponseBody(t *testing.T) {
var lines []string
logger := newRedactingLogger(captureLogf(&lines))
body := `{"key":"BASE64KEYMATERIAL==","keyVersion":3}`
logger.Debug("Body (length=%d): %s", len(body), body)
if strings.Contains(lines[0], "BASE64KEYMATERIAL") {
t.Fatalf("log line leaked key material: %q", lines[0])
}
}
func TestRedactingLoggerHidesByteSlices(t *testing.T) {
var lines []string
logger := newRedactingLogger(captureLogf(&lines))
logger.Info("payload=%s", []byte("plaintext-message"))
if strings.Contains(lines[0], "plaintext-message") {
t.Fatalf("log line leaked a byte payload: %q", lines[0])
}
}
func TestRedactingLoggerHidesErrorText(t *testing.T) {
var lines []string
logger := newRedactingLogger(captureLogf(&lines))
// Vendor error strings can embed a response body preview.
logger.Error("request failed: %v", errors.New(`server said {"key":"LEAKED"}`))
if strings.Contains(lines[0], "LEAKED") {
t.Fatalf("log line leaked error contents: %q", lines[0])
}
}
func TestRedactingLoggerKeepsDiagnosticNumbers(t *testing.T) {
var lines []string
logger := newRedactingLogger(captureLogf(&lines))
logger.Debug("Status: %d (took %s), size=%d", 503, 1500*time.Millisecond, 4096)
line := lines[0]
for _, want := range []string{"503", "1.5s", "4096"} {
if !strings.Contains(line, want) {
t.Fatalf("log line = %q, want it to keep %q for diagnostics", line, want)
}
}
}
func TestRedactingLoggerLabelsLevel(t *testing.T) {
var lines []string
logger := newRedactingLogger(captureLogf(&lines))
logger.Debug("d")
logger.Info("i")
logger.Error("e")
if len(lines) != 3 {
t.Fatalf("got %d lines, want 3", len(lines))
}
for i, want := range []string{"debug", "info", "error"} {
if !strings.Contains(lines[i], want) {
t.Fatalf("line %d = %q, want level %q", i, lines[i], want)
}
}
}
func TestRedactingLoggerToleratesNilSinkAtEmit(t *testing.T) {
// Guard against a partially constructed logger being used.
var logger *redactingLogger
logger.Debug("must not panic %s", "value")
}
func TestRedactArgKeepsDurations(t *testing.T) {
if got := redactArg(2 * time.Second); got != "2s" {
t.Fatalf("redactArg(2s) = %v, want 2s", got)
}
}
func TestRedactArgElidesStrings(t *testing.T) {
got, ok := redactArg("secret").(string)
if !ok {
t.Fatalf("redactArg returned %T, want string", got)
}
if strings.Contains(got, "secret") {
t.Fatalf("redactArg = %q, want the value elided", got)
}
if !strings.Contains(got, "len=6") {
t.Fatalf("redactArg = %q, want the length preserved", got)
}
}
func TestRedactArgPassesThroughNumbers(t *testing.T) {
if got := redactArg(42); got != 42 {
t.Fatalf("redactArg(42) = %v, want 42", got)
}
if got := redactArg(true); got != true {
t.Fatalf("redactArg(true) = %v, want true", got)
}
}
+297
View File
@@ -0,0 +1,297 @@
// Copyright 2026 Alibaba Group
// Licensed under the Apache License, Version 2.0 (the "License");
// you may not use this file except in compliance with the License.
// You may obtain a copy of the License at
//
// http://www.apache.org/licenses/LICENSE-2.0
//
// Unless required by applicable law or agreed to in writing, software
// distributed under the License is distributed on an "AS IS" BASIS,
// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
// See the License for the specific language governing permissions and
// limitations under the License.
// Package msgcrypto wraps the third-party SafeChat SDK so DWS can decrypt
// DingTalk messages that an organization has encrypted with its own key
// material, and encrypt outbound ones.
//
// The SafeChat backend links a prebuilt C static library and therefore needs
// CGO. Because DWS ships CGO-free cross-compiled release binaries, the backend
// is compiled only under the "safechat" build tag:
//
// CGO_ENABLED=1 go build -tags safechat ./cmd
//
// Every other build gets a stub whose constructor fails with ErrUnavailable,
// so callers must always handle that error rather than assume the capability
// exists. Use Available to branch before offering the feature to a user.
//
// Key material is fetched from the vendor key server on demand, which requires
// a DingTalk 免登 authCode supplied through AuthCodeProvider. DWS does not mint
// that code itself; the caller injects a provider.
package msgcrypto
import (
"context"
"errors"
"fmt"
"io/fs"
"os"
"path/filepath"
"sync"
"time"
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/pkg/config"
)
// Errors reported by this package. Callers are expected to test for
// ErrUnavailable explicitly, because it is the normal outcome on every build
// that does not enable the safechat tag.
var (
// ErrUnavailable means this binary was built without the SafeChat
// backend, or for a platform the vendor does not ship a static library
// for (notably windows/arm64).
ErrUnavailable = errors.New("msgcrypto: SafeChat backend not built into this binary")
// ErrAlreadyOpen means a Cipher is already open. The underlying C
// library keeps global state, so only one may exist per process.
ErrAlreadyOpen = errors.New("msgcrypto: a cipher is already open in this process")
// ErrClosed is returned by a Cipher whose Close has already run.
ErrClosed = errors.New("msgcrypto: cipher is closed")
// ErrNoAuthCodeProvider means Config.AuthCode was nil. Without it the
// backend cannot fetch or rotate key material.
ErrNoAuthCodeProvider = errors.New("msgcrypto: config.AuthCode is required")
// ErrEmptyPayload means an encrypt or decrypt call got no bytes. The
// vendor SDK rejects empty input, so we reject it earlier with a
// clearer message.
ErrEmptyPayload = errors.New("msgcrypto: payload is empty")
// ErrNoCorpID means the caller omitted the organization id, which
// selects the key and therefore cannot be defaulted.
ErrNoCorpID = errors.New("msgcrypto: corpID is required")
// ErrNoKeyServer means Config.KeyServer was empty. The vendor C
// library would otherwise pick the key-request destination.
ErrNoKeyServer = errors.New("msgcrypto: config.KeyServer is required")
// ErrInvalidKeyServer means Config.KeyServer is not a usable URL.
ErrInvalidKeyServer = errors.New("msgcrypto: config.KeyServer is not a valid URL")
// ErrKeyServerNotHTTPS means Config.KeyServer is not https.
ErrKeyServerNotHTTPS = errors.New("msgcrypto: config.KeyServer must be an https URL")
// ErrRedirectHostMismatch means the domain goProxy received does not
// match Config.AllowedRedirectHost. The domain is never sent to portal.
ErrRedirectHostMismatch = errors.New("msgcrypto: goProxy domain host does not match AllowedRedirectHost")
)
// keystoreDirPerm keeps the key cache owner-only. The directory holds
// organization key material, so it must not be group- or world-readable.
const keystoreDirPerm fs.FileMode = 0o700
// DefaultKeystoreDir returns the default key cache directory,
// ~/.dws/safechat/keystore, honouring DWS_CONFIG_DIR like the rest of DWS.
func DefaultKeystoreDir() string {
return filepath.Join(config.DefaultConfigDir(), "safechat", "keystore")
}
// Cipher encrypts and decrypts message payloads for one organization at a
// time. Implementations are safe for concurrent use.
type Cipher interface {
// EncryptMessage encrypts plaintext for corpID/staffID and returns the
// vendor ciphertext envelope.
EncryptMessage(ctx context.Context, corpID, staffID string, plaintext []byte) ([]byte, error)
// DecryptMessage decrypts a vendor ciphertext envelope.
DecryptMessage(ctx context.Context, corpID, staffID string, ciphertext []byte) ([]byte, error)
// Close releases the backend and frees the process-wide slot so a later
// Open can succeed. Calling it twice is safe.
Close() error
}
// Config parameterises Open.
type Config struct {
// KeystoreDir is where fetched keys are cached. Defaults to
// DefaultKeystoreDir. It is created with 0700 if missing.
KeystoreDir string
// UserID is an opaque local identifier. The vendor SDK stores it but
// does not use it for key selection; leave empty to let the SDK
// generate one.
UserID string
// AuthCode supplies the DingTalk 免登 authCode used to authenticate key
// requests. Required. The backend calls it only from the vendor goProxy
// callback (cold keystore or key-version rotation), never on every
// encrypt/decrypt.
AuthCode AuthCodeProvider
// KeyServer is the HTTPS URL of the vendor key service. Required: it
// replaces the host the closed-source C library would otherwise pick.
KeyServer string
// AllowedRedirectHost, when set, is compared to the domain the C
// library passes into goProxy. A mismatch fails that key fetch. It is
// a local check only; the domain is never sent to portal.
AllowedRedirectHost string
// MaxRetry bounds retries while a key is still being fetched. Zero
// selects the vendor default.
MaxRetry int
// HTTPTimeout bounds each key request. Zero selects the vendor default.
HTTPTimeout time.Duration
// Debug enables backend logging through a redacting logger. Off by
// default because the vendor SDK logs the authCode and raw key-server
// responses at debug level.
Debug bool
// Logf receives already-redacted backend log lines when Debug is set.
// Nil discards them.
Logf func(format string, args ...any)
}
// withDefaults returns cfg with empty optional fields filled in.
func (cfg Config) withDefaults() Config {
if cfg.KeystoreDir == "" {
cfg.KeystoreDir = DefaultKeystoreDir()
}
return cfg
}
// validate reports whether cfg carries everything the backend needs.
func (cfg Config) validate() error {
if cfg.AuthCode == nil {
return ErrNoAuthCodeProvider
}
if cfg.KeystoreDir == "" {
return errors.New("msgcrypto: config.KeystoreDir resolved to an empty path")
}
if err := validateKeyServer(cfg.KeyServer); err != nil {
return err
}
return nil
}
// prepareKeystore creates dir if needed and makes sure it is owner-only.
// An existing directory with looser bits is tightened, because it caches key
// material.
func prepareKeystore(dir string) error {
if err := os.MkdirAll(dir, keystoreDirPerm); err != nil {
return fmt.Errorf("msgcrypto: create keystore dir: %w", err)
}
info, err := os.Stat(dir)
if err != nil {
return fmt.Errorf("msgcrypto: stat keystore dir: %w", err)
}
if !info.IsDir() {
return fmt.Errorf("msgcrypto: keystore path %s is not a directory", dir)
}
// Windows does not model POSIX bits, so only tighten where they apply.
if runtimeSupportsPOSIXPerm && info.Mode().Perm() != keystoreDirPerm {
if err := os.Chmod(dir, keystoreDirPerm); err != nil {
return fmt.Errorf("msgcrypto: restrict keystore dir permissions: %w", err)
}
}
return nil
}
// process holds the single-instance guard. The vendor C library keeps global
// state, so a second concurrent Cipher would corrupt it.
var process struct {
mu sync.Mutex
open bool
}
// Open validates cfg, prepares the keystore and starts the backend.
//
// It returns ErrUnavailable when the backend was not compiled in, so callers
// can degrade gracefully. Only one Cipher may be open per process; Close frees
// the slot.
func Open(ctx context.Context, cfg Config) (Cipher, error) {
cfg = cfg.withDefaults()
if err := cfg.validate(); err != nil {
return nil, err
}
if !Available() {
return nil, ErrUnavailable
}
if err := prepareKeystore(cfg.KeystoreDir); err != nil {
return nil, err
}
process.mu.Lock()
defer process.mu.Unlock()
if process.open {
return nil, ErrAlreadyOpen
}
backend, err := newBackend(ctx, cfg)
if err != nil {
return nil, err
}
process.open = true
return &trackedCipher{backend: backend}, nil
}
// trackedCipher releases the process-wide slot when the wrapped backend closes.
type trackedCipher struct {
mu sync.Mutex
backend Cipher
}
// EncryptMessage validates the payload and delegates to the backend.
func (c *trackedCipher) EncryptMessage(ctx context.Context, corpID, staffID string, plaintext []byte) ([]byte, error) {
backend, err := c.live(corpID, plaintext)
if err != nil {
return nil, err
}
return backend.EncryptMessage(ctx, corpID, staffID, plaintext)
}
// DecryptMessage validates the payload and delegates to the backend.
func (c *trackedCipher) DecryptMessage(ctx context.Context, corpID, staffID string, ciphertext []byte) ([]byte, error) {
backend, err := c.live(corpID, ciphertext)
if err != nil {
return nil, err
}
return backend.DecryptMessage(ctx, corpID, staffID, ciphertext)
}
// live returns the backend after checking the cipher is open and the arguments
// are usable.
func (c *trackedCipher) live(corpID string, payload []byte) (Cipher, error) {
if corpID == "" {
return nil, ErrNoCorpID
}
if len(payload) == 0 {
return nil, ErrEmptyPayload
}
c.mu.Lock()
defer c.mu.Unlock()
if c.backend == nil {
return nil, ErrClosed
}
return c.backend, nil
}
// Close closes the backend once and releases the process-wide slot.
func (c *trackedCipher) Close() error {
c.mu.Lock()
backend := c.backend
c.backend = nil
c.mu.Unlock()
if backend == nil {
return nil
}
err := backend.Close()
process.mu.Lock()
process.open = false
process.mu.Unlock()
return err
}
+345
View File
@@ -0,0 +1,345 @@
// Copyright 2026 Alibaba Group
// Licensed under the Apache License, Version 2.0 (the "License");
// you may not use this file except in compliance with the License.
// You may obtain a copy of the License at
//
// http://www.apache.org/licenses/LICENSE-2.0
//
// Unless required by applicable law or agreed to in writing, software
// distributed under the License is distributed on an "AS IS" BASIS,
// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
// See the License for the specific language governing permissions and
// limitations under the License.
package msgcrypto
import (
"context"
"errors"
"io/fs"
"os"
"path/filepath"
"runtime"
"strings"
"sync"
"testing"
)
// fakeCipher stands in for a backend so the wrapper logic can be tested on
// every platform, with or without the safechat tag.
type fakeCipher struct {
mu sync.Mutex
encrypted [][]byte
decrypted [][]byte
closeCount int
closeErr error
}
func (f *fakeCipher) EncryptMessage(_ context.Context, _, _ string, plaintext []byte) ([]byte, error) {
f.mu.Lock()
defer f.mu.Unlock()
f.encrypted = append(f.encrypted, plaintext)
return []byte("cipher:" + string(plaintext)), nil
}
func (f *fakeCipher) DecryptMessage(_ context.Context, _, _ string, ciphertext []byte) ([]byte, error) {
f.mu.Lock()
defer f.mu.Unlock()
f.decrypted = append(f.decrypted, ciphertext)
return []byte(strings.TrimPrefix(string(ciphertext), "cipher:")), nil
}
func (f *fakeCipher) Close() error {
f.mu.Lock()
defer f.mu.Unlock()
f.closeCount++
return f.closeErr
}
// newTrackedForTest wraps backend and claims the process slot the same way
// Open does, so slot release can be asserted.
func newTrackedForTest(t *testing.T, backend Cipher) *trackedCipher {
t.Helper()
process.mu.Lock()
process.open = true
process.mu.Unlock()
t.Cleanup(func() {
process.mu.Lock()
process.open = false
process.mu.Unlock()
})
return &trackedCipher{backend: backend}
}
func TestConfigWithDefaultsFillsKeystoreDir(t *testing.T) {
cfg := Config{}.withDefaults()
if cfg.KeystoreDir == "" {
t.Fatal("withDefaults left KeystoreDir empty")
}
if want := DefaultKeystoreDir(); cfg.KeystoreDir != want {
t.Fatalf("KeystoreDir = %q, want %q", cfg.KeystoreDir, want)
}
}
func TestConfigWithDefaultsKeepsExplicitKeystoreDir(t *testing.T) {
cfg := Config{KeystoreDir: "/custom/keys"}.withDefaults()
if cfg.KeystoreDir != "/custom/keys" {
t.Fatalf("KeystoreDir = %q, want /custom/keys", cfg.KeystoreDir)
}
}
func TestDefaultKeystoreDirHonoursConfigDirOverride(t *testing.T) {
t.Setenv("DWS_CONFIG_DIR", t.TempDir())
got := DefaultKeystoreDir()
if !strings.HasSuffix(got, filepath.Join("safechat", "keystore")) {
t.Fatalf("DefaultKeystoreDir() = %q, want it to end with safechat/keystore", got)
}
if !strings.HasPrefix(got, os.Getenv("DWS_CONFIG_DIR")) {
t.Fatalf("DefaultKeystoreDir() = %q, want it under DWS_CONFIG_DIR", got)
}
}
func TestConfigValidateRequiresAuthCodeProvider(t *testing.T) {
err := Config{KeystoreDir: "/tmp/keys"}.validate()
if !errors.Is(err, ErrNoAuthCodeProvider) {
t.Fatalf("validate() = %v, want ErrNoAuthCodeProvider", err)
}
}
func TestConfigValidateAcceptsCompleteConfig(t *testing.T) {
cfg := Config{
KeystoreDir: "/tmp/keys",
AuthCode: StaticAuthCode("code"),
KeyServer: "https://key.example.test",
}
if err := cfg.validate(); err != nil {
t.Fatalf("validate() = %v, want nil", err)
}
}
func TestConfigValidateRequiresKeyServer(t *testing.T) {
err := Config{KeystoreDir: "/tmp/keys", AuthCode: StaticAuthCode("code")}.validate()
if !errors.Is(err, ErrNoKeyServer) {
t.Fatalf("validate() = %v, want ErrNoKeyServer", err)
}
}
func TestConfigValidateRejectsHTTPKeyServer(t *testing.T) {
err := Config{
KeystoreDir: "/tmp/keys",
AuthCode: StaticAuthCode("code"),
KeyServer: "http://key.example.test",
}.validate()
if !errors.Is(err, ErrKeyServerNotHTTPS) {
t.Fatalf("validate() = %v, want ErrKeyServerNotHTTPS", err)
}
}
func TestPrepareKeystoreCreatesOwnerOnlyDir(t *testing.T) {
dir := filepath.Join(t.TempDir(), "nested", "keystore")
if err := prepareKeystore(dir); err != nil {
t.Fatalf("prepareKeystore() = %v", err)
}
info, err := os.Stat(dir)
if err != nil {
t.Fatalf("stat: %v", err)
}
if !info.IsDir() {
t.Fatal("prepareKeystore did not create a directory")
}
if runtime.GOOS == "windows" {
return
}
if got := info.Mode().Perm(); got != keystoreDirPerm {
t.Fatalf("perm = %#o, want %#o", got, keystoreDirPerm)
}
}
func TestPrepareKeystoreTightensLoosePermissions(t *testing.T) {
if runtime.GOOS == "windows" {
t.Skip("POSIX permission bits are not modelled on Windows")
}
dir := filepath.Join(t.TempDir(), "keystore")
if err := os.MkdirAll(dir, fs.FileMode(0o755)); err != nil {
t.Fatalf("mkdir: %v", err)
}
if err := prepareKeystore(dir); err != nil {
t.Fatalf("prepareKeystore() = %v", err)
}
info, err := os.Stat(dir)
if err != nil {
t.Fatalf("stat: %v", err)
}
if got := info.Mode().Perm(); got != keystoreDirPerm {
t.Fatalf("perm = %#o, want %#o (key material must stay owner-only)", got, keystoreDirPerm)
}
}
func TestPrepareKeystoreRejectsFilePath(t *testing.T) {
path := filepath.Join(t.TempDir(), "keystore")
if err := os.WriteFile(path, []byte("not a dir"), 0o600); err != nil {
t.Fatalf("write: %v", err)
}
err := prepareKeystore(path)
if err == nil {
t.Fatal("prepareKeystore() = nil, want an error for a non-directory path")
}
}
func TestOpenRejectsMissingAuthCodeProviderBeforeTouchingDisk(t *testing.T) {
dir := filepath.Join(t.TempDir(), "keystore")
_, err := Open(context.Background(), Config{KeystoreDir: dir})
if !errors.Is(err, ErrNoAuthCodeProvider) {
t.Fatalf("Open() = %v, want ErrNoAuthCodeProvider", err)
}
if _, statErr := os.Stat(dir); !os.IsNotExist(statErr) {
t.Fatal("Open created the keystore dir despite an invalid config")
}
}
func TestOpenRejectsMissingKeyServerBeforeTouchingDisk(t *testing.T) {
dir := filepath.Join(t.TempDir(), "keystore")
_, err := Open(context.Background(), Config{
KeystoreDir: dir,
AuthCode: StaticAuthCode("code"),
})
if !errors.Is(err, ErrNoKeyServer) {
t.Fatalf("Open() = %v, want ErrNoKeyServer", err)
}
if _, statErr := os.Stat(dir); !os.IsNotExist(statErr) {
t.Fatal("Open created the keystore dir despite a missing KeyServer")
}
}
func TestOpenWithoutBackendReportsUnavailable(t *testing.T) {
if Available() {
t.Skip("this binary has the safechat backend compiled in")
}
_, err := Open(context.Background(), Config{
KeystoreDir: filepath.Join(t.TempDir(), "keystore"),
AuthCode: StaticAuthCode("code"),
KeyServer: "https://key.example.test",
})
if !errors.Is(err, ErrUnavailable) {
t.Fatalf("Open() = %v, want ErrUnavailable", err)
}
}
func TestAvailableAgreesWithBackendVersion(t *testing.T) {
if Available() != (BackendVersion != "") {
t.Fatalf("Available() = %v but BackendVersion = %q; they must agree", Available(), BackendVersion)
}
}
func TestTrackedCipherRoundTripsThroughBackend(t *testing.T) {
backend := &fakeCipher{}
cipher := newTrackedForTest(t, backend)
ciphertext, err := cipher.EncryptMessage(context.Background(), "corp", "staff", []byte("hello"))
if err != nil {
t.Fatalf("EncryptMessage() = %v", err)
}
plaintext, err := cipher.DecryptMessage(context.Background(), "corp", "staff", ciphertext)
if err != nil {
t.Fatalf("DecryptMessage() = %v", err)
}
if string(plaintext) != "hello" {
t.Fatalf("round trip = %q, want hello", plaintext)
}
}
func TestTrackedCipherRejectsEmptyCorpID(t *testing.T) {
cipher := newTrackedForTest(t, &fakeCipher{})
if _, err := cipher.EncryptMessage(context.Background(), "", "staff", []byte("x")); !errors.Is(err, ErrNoCorpID) {
t.Fatalf("EncryptMessage() = %v, want ErrNoCorpID", err)
}
if _, err := cipher.DecryptMessage(context.Background(), "", "staff", []byte("x")); !errors.Is(err, ErrNoCorpID) {
t.Fatalf("DecryptMessage() = %v, want ErrNoCorpID", err)
}
}
func TestTrackedCipherRejectsEmptyPayload(t *testing.T) {
cipher := newTrackedForTest(t, &fakeCipher{})
if _, err := cipher.EncryptMessage(context.Background(), "corp", "staff", nil); !errors.Is(err, ErrEmptyPayload) {
t.Fatalf("EncryptMessage() = %v, want ErrEmptyPayload", err)
}
if _, err := cipher.DecryptMessage(context.Background(), "corp", "staff", []byte{}); !errors.Is(err, ErrEmptyPayload) {
t.Fatalf("DecryptMessage() = %v, want ErrEmptyPayload", err)
}
}
func TestTrackedCipherRejectsUseAfterClose(t *testing.T) {
backend := &fakeCipher{}
cipher := newTrackedForTest(t, backend)
if err := cipher.Close(); err != nil {
t.Fatalf("Close() = %v", err)
}
if _, err := cipher.EncryptMessage(context.Background(), "corp", "staff", []byte("x")); !errors.Is(err, ErrClosed) {
t.Fatalf("EncryptMessage() after Close = %v, want ErrClosed", err)
}
}
func TestTrackedCipherCloseIsIdempotentAndClosesBackendOnce(t *testing.T) {
backend := &fakeCipher{}
cipher := newTrackedForTest(t, backend)
for i := 0; i < 3; i++ {
if err := cipher.Close(); err != nil {
t.Fatalf("Close() #%d = %v", i+1, err)
}
}
if backend.closeCount != 1 {
t.Fatalf("backend Close called %d times, want exactly 1", backend.closeCount)
}
}
func TestTrackedCipherCloseReleasesProcessSlot(t *testing.T) {
cipher := newTrackedForTest(t, &fakeCipher{})
if err := cipher.Close(); err != nil {
t.Fatalf("Close() = %v", err)
}
process.mu.Lock()
open := process.open
process.mu.Unlock()
if open {
t.Fatal("Close did not release the process slot, so a later Open would fail")
}
}
func TestTrackedCipherClosePropagatesBackendError(t *testing.T) {
wantErr := errors.New("backend close failed")
cipher := newTrackedForTest(t, &fakeCipher{closeErr: wantErr})
if err := cipher.Close(); !errors.Is(err, wantErr) {
t.Fatalf("Close() = %v, want %v", err, wantErr)
}
process.mu.Lock()
open := process.open
process.mu.Unlock()
if open {
t.Fatal("a failing backend Close must still release the process slot")
}
}
func TestOpenRefusesSecondCipherWhileOneIsOpen(t *testing.T) {
// The vendor C library keeps global state, so a second concurrent
// cipher must be refused rather than silently corrupt it.
process.mu.Lock()
process.open = true
process.mu.Unlock()
t.Cleanup(func() {
process.mu.Lock()
process.open = false
process.mu.Unlock()
})
if !Available() {
t.Skip("Open reports ErrUnavailable before reaching the single-instance guard")
}
_, err := Open(context.Background(), Config{
KeystoreDir: filepath.Join(t.TempDir(), "keystore"),
AuthCode: StaticAuthCode("code"),
KeyServer: "https://key.example.test",
})
if !errors.Is(err, ErrAlreadyOpen) {
t.Fatalf("Open() = %v, want ErrAlreadyOpen", err)
}
}
+20
View File
@@ -0,0 +1,20 @@
// Copyright 2026 Alibaba Group
// Licensed under the Apache License, Version 2.0 (the "License");
// you may not use this file except in compliance with the License.
// You may obtain a copy of the License at
//
// http://www.apache.org/licenses/LICENSE-2.0
//
// Unless required by applicable law or agreed to in writing, software
// distributed under the License is distributed on an "AS IS" BASIS,
// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
// See the License for the specific language governing permissions and
// limitations under the License.
//go:build !windows
package msgcrypto
// runtimeSupportsPOSIXPerm reports that the keystore directory's permission
// bits are meaningful here, so prepareKeystore enforces owner-only access.
const runtimeSupportsPOSIXPerm = true
+21
View File
@@ -0,0 +1,21 @@
// Copyright 2026 Alibaba Group
// Licensed under the Apache License, Version 2.0 (the "License");
// you may not use this file except in compliance with the License.
// You may obtain a copy of the License at
//
// http://www.apache.org/licenses/LICENSE-2.0
//
// Unless required by applicable law or agreed to in writing, software
// distributed under the License is distributed on an "AS IS" BASIS,
// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
// See the License for the specific language governing permissions and
// limitations under the License.
//go:build windows
package msgcrypto
// runtimeSupportsPOSIXPerm reports that Windows does not model POSIX
// permission bits, so prepareKeystore leaves the directory mode alone and
// relies on the user profile ACL instead.
const runtimeSupportsPOSIXPerm = false
+170
View File
@@ -0,0 +1,170 @@
// Copyright 2026 Alibaba Group
// Licensed under the Apache License, Version 2.0 (the "License");
// you may not use this file except in compliance with the License.
// You may obtain a copy of the License at
//
// http://www.apache.org/licenses/LICENSE-2.0
//
// Unless required by applicable law or agreed to in writing, software
// distributed under the License is distributed on an "AS IS" BASIS,
// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
// See the License for the specific language governing permissions and
// limitations under the License.
package msgcrypto
import (
"context"
"errors"
"fmt"
"net/http"
"strings"
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/auth"
)
// VendorSafeChat is the first vendorAuthCode vendor.
const VendorSafeChat = "safechat"
// PortalAuthCode mints a one-shot 免登 authCode from portal
// POST /oauth2/vendorAuthCode. It does not cache the code: goProxy spends it
// immediately. Do not wrap this in CachedAuthCode.
type PortalAuthCode struct {
ConfigDir string
Vendor string
CLIVersion string
HTTPClient *http.Client
clientID func() string
snapshot func(ctx context.Context, configDir string) (*auth.TokenData, error)
refresh func(ctx context.Context, configDir, rejected string) (string, error)
fetch func(ctx context.Context, in auth.VendorAuthCodeInput) (*auth.VendorAuthCodeResult, error)
}
// NewPortalAuthCode returns a provider that talks to portal with the current
// login. cliVersion is sent as x-dws-cli-version; leave empty only when the
// caller cannot know the CLI version.
func NewPortalAuthCode(configDir, cliVersion string) *PortalAuthCode {
return &PortalAuthCode{
ConfigDir: configDir,
Vendor: VendorSafeChat,
CLIVersion: cliVersion,
clientID: auth.ClientID,
snapshot: func(ctx context.Context, dir string) (*auth.TokenData, error) {
return auth.NewOAuthProvider(dir, nil).GetTokenSnapshot(ctx)
},
refresh: func(ctx context.Context, dir, rejected string) (string, error) {
return auth.NewOAuthProvider(dir, nil).ForceRefreshRejectedToken(ctx, rejected)
},
fetch: auth.FetchVendorAuthCode,
}
}
// AuthCode mints a code for the logged-in organization. Prefer AuthCodeForCorp
// when goProxy already has a corpID.
func (p *PortalAuthCode) AuthCode(ctx context.Context) (string, error) {
snap, err := p.loadSnapshot(ctx)
if err != nil {
return "", err
}
return p.AuthCodeForCorp(ctx, snap.CorpID)
}
// AuthCodeForCorp mints a one-shot code for corpID. The request body is only
// vendor + corpId.
func (p *PortalAuthCode) AuthCodeForCorp(ctx context.Context, corpID string) (string, error) {
corpID = strings.TrimSpace(corpID)
if corpID == "" {
return "", ErrNoCorpID
}
snap, err := p.loadSnapshot(ctx)
if err != nil {
return "", err
}
token := strings.TrimSpace(snap.AccessToken)
if token == "" {
return "", fmt.Errorf("msgcrypto: access token is empty")
}
result, err := p.fetchOnce(ctx, snap, token, corpID)
if err == nil {
return result.AuthCode, nil
}
var verr *auth.VendorAuthCodeError
if errors.As(err, &verr) && verr.Code == auth.VendorAuthCodeTokenInvalid {
refreshed, rerr := p.doRefresh(ctx, token)
if rerr == nil && refreshed != "" && refreshed != token {
result, err = p.fetchOnce(ctx, snap, refreshed, corpID)
if err == nil {
return result.AuthCode, nil
}
}
return "", err
}
if retryableVendorAuthCode(err) {
result, err = p.fetchOnce(ctx, snap, token, corpID)
if err == nil {
return result.AuthCode, nil
}
}
return "", err
}
func (p *PortalAuthCode) loadSnapshot(ctx context.Context) (*auth.TokenData, error) {
if p.snapshot == nil {
return nil, fmt.Errorf("msgcrypto: portal authCode snapshot loader is not configured")
}
snap, err := p.snapshot(ctx, p.ConfigDir)
if err != nil {
return nil, fmt.Errorf("msgcrypto: load access token: %w", err)
}
if snap == nil {
return nil, fmt.Errorf("msgcrypto: load access token: empty snapshot")
}
return snap, nil
}
func (p *PortalAuthCode) fetchOnce(ctx context.Context, snap *auth.TokenData, token, corpID string) (*auth.VendorAuthCodeResult, error) {
fetch := p.fetch
if fetch == nil {
fetch = auth.FetchVendorAuthCode
}
vendor := strings.TrimSpace(p.Vendor)
if vendor == "" {
vendor = VendorSafeChat
}
clientID := strings.TrimSpace(snap.ClientID)
if clientID == "" && p.clientID != nil {
clientID = strings.TrimSpace(p.clientID())
}
return fetch(ctx, auth.VendorAuthCodeInput{
AccessToken: token,
ClientID: clientID,
CLIVersion: p.CLIVersion,
LoginRegion: auth.LoginRegion(strings.TrimSpace(snap.LoginRegion)),
Vendor: vendor,
CorpID: corpID,
HTTPClient: p.HTTPClient,
})
}
func (p *PortalAuthCode) doRefresh(ctx context.Context, rejected string) (string, error) {
if p.refresh == nil {
return "", fmt.Errorf("msgcrypto: token refresh is not configured")
}
return p.refresh(ctx, p.ConfigDir, rejected)
}
func retryableVendorAuthCode(err error) bool {
var verr *auth.VendorAuthCodeError
if errors.As(err, &verr) {
return verr.Retryable()
}
var statusErr *auth.HTTPStatusError
if errors.As(err, &statusErr) && statusErr != nil {
return statusErr.StatusCode == http.StatusTooManyRequests ||
statusErr.StatusCode >= http.StatusInternalServerError
}
return false
}
+121
View File
@@ -0,0 +1,121 @@
// Copyright 2026 Alibaba Group
// Licensed under the Apache License, Version 2.0 (the "License");
// you may not use this file except in compliance with the License.
// You may obtain a copy of the License at
//
// http://www.apache.org/licenses/LICENSE-2.0
//
// Unless required by applicable law or agreed to in writing, software
// distributed under the License is distributed on an "AS IS" BASIS,
// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
// See the License for the specific language governing permissions and
// limitations under the License.
package msgcrypto
import (
"context"
"errors"
"sync/atomic"
"testing"
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/auth"
)
func TestPortalAuthCodePostsVendorAndCorpOnly(t *testing.T) {
var got auth.VendorAuthCodeInput
p := &PortalAuthCode{
ConfigDir: t.TempDir(),
Vendor: VendorSafeChat,
CLIVersion: "1.2.3",
snapshot: func(context.Context, string) (*auth.TokenData, error) {
return &auth.TokenData{
AccessToken: "user-token",
ClientID: "dws-client",
CorpID: "ding_login",
LoginRegion: "",
}, nil
},
fetch: func(_ context.Context, in auth.VendorAuthCodeInput) (*auth.VendorAuthCodeResult, error) {
got = in
return &auth.VendorAuthCodeResult{AuthCode: "tmp-code", ExpiresIn: 120}, nil
},
}
code, err := p.AuthCodeForCorp(context.Background(), "ding_target")
if err != nil || code != "tmp-code" {
t.Fatalf("AuthCodeForCorp() = %q, %v", code, err)
}
if got.Vendor != VendorSafeChat || got.CorpID != "ding_target" {
t.Fatalf("body vendor/corpId = %q/%q", got.Vendor, got.CorpID)
}
if got.AccessToken != "user-token" || got.ClientID != "dws-client" || got.CLIVersion != "1.2.3" {
t.Fatalf("headers token/client/version = %q/%q/%q", got.AccessToken, got.ClientID, got.CLIVersion)
}
}
func TestPortalAuthCodeRetriesOnceOnTokenInvalidAfterRefresh(t *testing.T) {
var calls atomic.Int32
p := &PortalAuthCode{
snapshot: func(context.Context, string) (*auth.TokenData, error) {
return &auth.TokenData{AccessToken: "stale", ClientID: "dws-client", CorpID: "ding"}, nil
},
refresh: func(context.Context, string, string) (string, error) {
return "fresh", nil
},
fetch: func(_ context.Context, in auth.VendorAuthCodeInput) (*auth.VendorAuthCodeResult, error) {
n := calls.Add(1)
if in.AccessToken == "stale" {
return nil, &auth.VendorAuthCodeError{Code: auth.VendorAuthCodeTokenInvalid, Message: "expired"}
}
if in.AccessToken != "fresh" {
t.Fatalf("retry token = %q, want fresh", in.AccessToken)
}
if n != 2 {
t.Fatalf("fetch calls = %d, want 2", n)
}
return &auth.VendorAuthCodeResult{AuthCode: "new-code", ExpiresIn: 120}, nil
},
}
code, err := p.AuthCodeForCorp(context.Background(), "ding")
if err != nil || code != "new-code" {
t.Fatalf("AuthCodeForCorp() = %q, %v", code, err)
}
if calls.Load() != 2 {
t.Fatalf("fetch called %d times, want 2", calls.Load())
}
}
func TestPortalAuthCodeDoesNotRetryOrgMismatch(t *testing.T) {
var calls atomic.Int32
want := &auth.VendorAuthCodeError{Code: auth.VendorAuthCodeOrgMismatch, Message: "mismatch"}
p := &PortalAuthCode{
snapshot: func(context.Context, string) (*auth.TokenData, error) {
return &auth.TokenData{AccessToken: "tok", ClientID: "dws-client", CorpID: "ding"}, nil
},
fetch: func(context.Context, auth.VendorAuthCodeInput) (*auth.VendorAuthCodeResult, error) {
calls.Add(1)
return nil, want
},
}
_, err := p.AuthCodeForCorp(context.Background(), "other")
if !errors.Is(err, want) {
t.Fatalf("AuthCodeForCorp() = %v, want %v", err, want)
}
if calls.Load() != 1 {
t.Fatalf("fetch called %d times, want 1", calls.Load())
}
}
func TestPortalAuthCodeRequiresCorpID(t *testing.T) {
p := &PortalAuthCode{
snapshot: func(context.Context, string) (*auth.TokenData, error) {
return &auth.TokenData{AccessToken: "tok", ClientID: "id"}, nil
},
}
if _, err := p.AuthCodeForCorp(context.Background(), " "); !errors.Is(err, ErrNoCorpID) {
t.Fatalf("AuthCodeForCorp() = %v, want ErrNoCorpID", err)
}
}
+148
View File
@@ -0,0 +1,148 @@
package safechat
/*
#include "csrc/safechat.h"
#include "goproxy_bridge.h"
#include <stdlib.h>
*/
import "C"
import (
"sync/atomic"
"unsafe"
)
var globalClient atomic.Value // stores *Client
// registerGlobalClient stores the client reference for CGO callbacks.
func registerGlobalClient(c *Client) {
globalClient.Store(c)
}
// getGlobalClient retrieves the active client from the atomic value.
// Returns nil if no client is registered.
func getGlobalClient() *Client {
v := globalClient.Load()
if v == nil {
return nil
}
return v.(*Client)
}
// goProxy is the CGO callback invoked by the C library when a key request needs to be sent.
//
//export goProxy
func goProxy(corpid, uid, domain, url, param, seqID *C.char) C.int {
client := getGlobalClient()
if client == nil {
return -1
}
// Copy C strings to Go values immediately to avoid dangling pointer issues.
// The C layer may free these buffers after goProxy returns.
goCorpID := C.GoString(corpid)
goDomain := C.GoString(domain)
goURL := C.GoString(url)
goParam := C.GoString(param)
if client.cfg.Logger != nil {
client.cfg.Logger.Debug("goProxy called: corpID=%s, url=%s", goCorpID, goURL)
client.cfg.Logger.Debug("goProxy input param (length=%d): %s", len(goParam), previewString(goParam, 1024))
}
code := client.cfg.Code
if client.cfg.AuthCodeHook != nil {
hookCode, hookErr := client.cfg.AuthCodeHook(goCorpID, goDomain)
if hookErr != nil {
if client.cfg.Logger != nil {
client.cfg.Logger.Error("goProxy: authCode hook failed for corp %s: %v", goCorpID, hookErr)
}
return -1
}
if hookCode == "" {
if client.cfg.Logger != nil {
client.cfg.Logger.Error("goProxy: authCode hook returned an empty code for corp %s", goCorpID)
}
return -1
}
code = hookCode
}
// Use kcMu to protect the HTTP request - NOT c.mu!
// c.mu is already held by the caller (EncryptMsg/DecryptMsg etc.)
client.kcMu.Lock()
resp, err := client.kc.doKeyRequest(goURL, goParam, code)
client.kcMu.Unlock()
if err != nil {
if client.cfg.Logger != nil {
client.cfg.Logger.Error("goProxy: key request failed for corp %s: %v", goCorpID, err)
}
return -1
}
if client.cfg.Logger != nil {
client.cfg.Logger.Debug("goProxy: HTTP request succeeded, feeding response to C setResponse (corpID=%s, response length=%d): %s",
goCorpID, len(resp), previewString(resp, 1024))
}
cCorpID := C.CString(goCorpID)
cResp := C.CString(resp)
defer C.free(unsafe.Pointer(cCorpID))
defer C.free(unsafe.Pointer(cResp))
ret := C.setResponse(
cCorpID,
cResp,
(C.block_crypto_func)(C.goBlockBridge),
)
if client.cfg.Logger != nil {
client.cfg.Logger.Debug("goProxy: C.setResponse returned %d for corpID=%s", ret, goCorpID)
}
if ret != 0 {
if client.cfg.Logger != nil {
client.cfg.Logger.Error("goProxy: C.setResponse FAILED for corpID=%s, ret=%d (key file may NOT have been generated)", goCorpID, ret)
}
return -1
}
return 0
}
// goBlock is called by C library when an enterprise key becomes restricted.
// This sets a flag that can be checked by the Go application layer.
//
//export goBlock
func goBlock(corpid *C.char) C.int {
client := getGlobalClient()
if client == nil {
return -1
}
goCorpID := C.GoString(corpid)
if client.cfg.Logger != nil {
client.cfg.Logger.Info("goBlock: enterprise %s key restricted", goCorpID)
}
// Store blocked status - application can check via IsBlocked()
client.blockedCorps.Store(goCorpID, true)
return 0
}
// goCancelBlock is called when an enterprise key restriction is lifted.
//
//export goCancelBlock
func goCancelBlock(corpid *C.char) C.int {
client := getGlobalClient()
if client == nil {
return -1
}
goCorpID := C.GoString(corpid)
if client.cfg.Logger != nil {
client.cfg.Logger.Info("goCancelBlock: enterprise %s key unblocked", goCorpID)
}
client.blockedCorps.Delete(goCorpID)
return 0
}
+9
View File
@@ -0,0 +1,9 @@
//go:build darwin && amd64
package safechat
/*
#cgo CFLAGS: -I${SRCDIR}/csrc -I${SRCDIR}/csrc/include -DDLL_EXPORT
#cgo LDFLAGS: -L${SRCDIR}/lib/darwin_amd64 -lsafechat -lpthread -ldl -lm -framework Security
*/
import "C"
+9
View File
@@ -0,0 +1,9 @@
//go:build darwin && arm64
package safechat
/*
#cgo CFLAGS: -I${SRCDIR}/csrc -I${SRCDIR}/csrc/include -DDLL_EXPORT
#cgo LDFLAGS: -L${SRCDIR}/lib/darwin_arm64 -lsafechat -lpthread -ldl -lm -framework Security
*/
import "C"
+9
View File
@@ -0,0 +1,9 @@
//go:build linux && amd64
package safechat
/*
#cgo CFLAGS: -I${SRCDIR}/csrc -I${SRCDIR}/csrc/include -DDLL_EXPORT -D_GNU_SOURCE
#cgo LDFLAGS: -L${SRCDIR}/lib/linux_amd64 -lsafechat -lpthread -ldl -lm
*/
import "C"
+9
View File
@@ -0,0 +1,9 @@
//go:build linux && arm64
package safechat
/*
#cgo CFLAGS: -I${SRCDIR}/csrc -I${SRCDIR}/csrc/include -DDLL_EXPORT -D_GNU_SOURCE
#cgo LDFLAGS: -L${SRCDIR}/lib/linux_arm64 -lsafechat -lpthread -ldl -lm
*/
import "C"
+9
View File
@@ -0,0 +1,9 @@
//go:build windows && amd64
package safechat
/*
#cgo CFLAGS: -I${SRCDIR}/csrc -I${SRCDIR}/csrc/include -DDLL_EXPORT -D_WIN32 -DWIN32
#cgo LDFLAGS: ${SRCDIR}/lib/windows_amd64/libsafechat.a -lws2_32 -lgdi32 -lcrypt32 -ladvapi32 -luser32
*/
import "C"
+9
View File
@@ -0,0 +1,9 @@
//go:build windows && arm64
package safechat
/*
#cgo CFLAGS: -I${SRCDIR}/csrc -I${SRCDIR}/csrc/include -DDLL_EXPORT -D_WIN32 -DWIN32
#cgo LDFLAGS: -L${SRCDIR}/lib/windows_arm64 -lsafechat -lws2_32 -lgdi32 -lcrypt32 -ladvapi32 -luser32
*/
import "C"
+88
View File
@@ -0,0 +1,88 @@
package safechat
import (
"time"
"github.com/google/uuid"
)
// Logger defines the logging interface for the SDK.
// Implementations should be safe for concurrent use.
type Logger interface {
// Debug logs a debug-level message
Debug(msg string, args ...interface{})
// Info logs an info-level message
Info(msg string, args ...interface{})
// Error logs an error-level message
Error(msg string, args ...interface{})
}
// Config holds configuration parameters for the SafeChat client.
type Config struct {
// DataPath is the directory path for storing keys and related data (required).
// The directory must exist and be writable.
DataPath string
// UserID is an optional user identifier. If empty, a random UUID is
// generated automatically. The C library stores this value but does not
// use it for key operations, so any non-empty string works.
UserID string
// Code is a DingTalk 免登 authCode used for key server authentication
// when AuthCodeHook is nil. Prefer AuthCodeHook: the code is one-shot
// and should be minted only when goProxy actually needs a key.
//
// Required (via Code or AuthCodeHook) when:
// - First-time key fetch (empty keystore)
// - Server-side key version rotation
//
// If keys are already cached AND the key version has not changed,
// neither field is used (no network request is made).
Code string
// AuthCodeHook is called from goProxy immediately before the key
// request. corpID and domain come from the C library callback.
// domain is the vendor SSO host and must not be forwarded as an
// authorize redirectURI; it is only for local host checks.
// If set, it replaces Code for that request and the returned value
// is not stored on the client.
AuthCodeHook func(corpID, domain string) (string, error)
// KeyServer is an optional override for the key server URL.
// If empty, the URL provided by the C library's goProxy callback will be used.
KeyServer string
// MaxRetry is the maximum number of retry attempts when key is not yet available.
// Default: 5 (private server redirect consumes 2 retries, so 5 provides sufficient margin)
MaxRetry int
// HTTPTimeout is the timeout for HTTP key requests.
// Default: 10s
HTTPTimeout time.Duration
// Logger is an optional logger instance.
// If nil, no logging will be performed.
Logger Logger
}
// defaultConfig returns a Config with sensible defaults applied.
func defaultConfig(cfg Config) Config {
if cfg.MaxRetry <= 0 {
cfg.MaxRetry = 5
}
if cfg.HTTPTimeout <= 0 {
cfg.HTTPTimeout = 10 * time.Second
}
if cfg.UserID == "" {
cfg.UserID = uuid.New().String()
}
return cfg
}
// validate checks that required config fields are set.
func (cfg *Config) validate() error {
if cfg.DataPath == "" {
return ErrConfigDataPathEmpty
}
return nil
}
+137
View File
@@ -0,0 +1,137 @@
#ifndef SAFECHAT_H
#define SAFECHAT_H
#ifdef WIN32
#ifndef DLL_EXPORT
#define CREATEDLL_API __declspec(dllimport)
#else
#define CREATEDLL_API __declspec(dllexport)
#endif
#else
#define CREATEDLL_API __attribute__((visibility("default")))
#endif
#include <stdint.h>
#ifdef __cplusplus
extern "C" {
#endif
typedef unsigned char byte;
// using namespace std;
// #pragma comment(lib, "safechat.lib")
#define DTMsgRetOK 0 /*operate return success*/
#define WARNING_RET_CODE 0x10000000
enum { RET_CODE_SUCCESS = 0, RET_CODE_ERROR, RET_CODE_WARNING };
// 获取返回值的类型,RET_CODE_WARNING | RET_CODE_ERROR| RET_CODE_SUCCESS
// #define GET_RET_CODETYPE(x) (x == 0) ? RET_CODE_SUCCESS : ((-x) & WARNING_RET_CODE == WARNING_RET_CODE ?
// RET_CODE_WARNING : RET_CODE_ERROR)
/*
打印日志回调函数
@param:ret_type 打印日志类型,RET_CODE_WARNING | RET_CODE_ERROR| RET_CODE_SUCCESS
@param:msg_detail 打印日志的详细信息,char * UTF-8格式
@return 备用,暂且返回NULL
*/
typedef void *(*log_func)(int ret_type, char *msg_detail);
/*
代理回调函数
@param: corpid 企业id
@param: uid 当前用户id
@param: domain 请求的域名
@param: url 请求url
@param: param 请求参数
@return: 返回执行结果
@param: seq_id 回调函数线程识别参数
*/
typedef int (*call_proxy_func)(char *corpid, char *uid, char *domain, char *url, char *param,
char *seq_id);
/*
设置密钥管控标记
@param: corpid 企业id
*/
typedef int (*block_crypto_func)(char *corpid);
/*
解除密钥管控标记
@param: corpid 企业id
*/
typedef int (*cancel_block_crypto_func)(char *corpid);
/*
初始化函数
@param: path 用户数据路径, 保存key等信息
@param: my_id 当前登W录用户的id
@return: 返回执行结果
*/
CREATEDLL_API int safechatInit(char *path, char *my_id);
/*
数据加密函数
@param: corp_id 企业id
@param: staffid 员工id
@param: data_buf 加密前数据
@param: data_len data_buf长度
@param: id 对方用户id, 如果群消息设置为NULL
@param: encrypt_buf 加密后的数据指针,需要外部函数释放
@param: ret_len 加密数据的返回长度
@param: proxy 代理请求回调函数
@return: 返回执行结果
@param: seq_id 回调函数线程识别参数
*/
CREATEDLL_API int encryptData(char *corp_id, char *staffid, unsigned char *data_buf, unsigned int data_len, char *id,
unsigned char **encrypt_buf, unsigned int *ret_len, char *seq_id, call_proxy_func proxy);
/*
数据解密函数
@param: corp_id 企业id
@param: staffid 员工id
@param: data_buf 解密前数据
@param: data_len data_buf长度
@param: id 对方用户id, 如果群消息设置为NULL
@param: decrypt_buf 解密后的数据指针,需要外部函数释放
@param: ret_len 解密数据的返回长度
@param: proxy 代理请求回调函数
@return: 返回执行结果
@param: seq_id 回调函数线程识别参数
*/
CREATEDLL_API int decryptData(char *corp_id, char *staffid, unsigned char *data_buf, unsigned int data_len, char *id,
unsigned char **decrypt_buf, unsigned int *ret_len, char *seq_id, call_proxy_func proxy);
CREATEDLL_API int encryptFile(char *corp_id, char *staffid, char *id, char *in_file_path, char *out_file_path,
char *seq_id, call_proxy_func proxy);
CREATEDLL_API int encryptBuffer(char *corp_id, char *staffid, byte *data_buf, uint32_t data_len, char *id,
byte **encrypt_buf, uint32_t *ret_len, char *seq_id, call_proxy_func proxy);
CREATEDLL_API int decryptFile(char *corp_id, char *staffid, char *id, char *in_file_path, char *out_file_path,
char *seq_id, call_proxy_func proxy);
CREATEDLL_API int decryptBuffer(char *corp_id, char *staffid, byte *data_buf, uint32_t data_len, char *id,
byte **decrypt_buf, uint32_t *ret_len, char *seq_id, call_proxy_func proxy);
/*
服务器响应处理
@param: corp_id 企业id
@param: json_str json格式的服务器返回值
@return: 返回执行结果
*/
CREATEDLL_API int setResponse(char *corp_id, char *json_str, block_crypto_func block_func1);
/*
消息推送处理
@param: corp_id 企业id
@param: staffid 员工id
@param: push_data 发送请求的数据
@param: proxy 代理请求回调函数
@return: 返回执行结果
@param: seq_id 回调函数线程识别参数
*/
CREATEDLL_API int setPushData(char *corp_id, char *staffid, char *push_data, char *seq_id, call_proxy_func proxy,
cancel_block_crypto_func cancel_crypto_func);
#ifdef __cplusplus
}
#endif
#endif // SAFECHAT_H
+265
View File
@@ -0,0 +1,265 @@
package safechat
import (
"errors"
"fmt"
)
// SDK-level errors (Go side)
var (
// ErrConfigDataPathEmpty indicates DataPath is not set in Config.
ErrConfigDataPathEmpty = errors.New("safechat: config.DataPath is required")
// ErrConfigUserIDEmpty indicates UserID is not set in Config.
ErrConfigUserIDEmpty = errors.New("safechat: config.UserID is required")
// ErrNotInitialized indicates the client has not been initialized.
ErrNotInitialized = errors.New("safechat: client not initialized")
// ErrMaxRetryExceeded indicates max retry attempts for key acquisition exceeded.
ErrMaxRetryExceeded = errors.New("safechat: max retry exceeded, key not available")
// ErrKeyRestricted indicates the enterprise key is restricted (managed/blocked).
ErrKeyRestricted = errors.New("safechat: enterprise key is restricted by admin")
// ErrKeyRequestFailed indicates the HTTP key request failed.
ErrKeyRequestFailed = errors.New("safechat: key request HTTP call failed")
// ErrAlreadyInitialized indicates init was called more than once.
ErrAlreadyInitialized = errors.New("safechat: client already initialized")
)
// CError represents an error code returned from the C library layer.
type CError struct {
Code int
Message string
}
func (e *CError) Error() string {
return fmt.Sprintf("safechat: C library error %d: %s", e.Code, e.Message)
}
// mapCError maps a C library return code to a Go error.
// Returns nil if code == 0 (FUNCTION_OK).
func mapCError(code int) error {
if code == 0 {
return nil
}
msg, ok := cErrorMessages[code]
if !ok {
msg = "unknown error"
}
return &CError{Code: code, Message: msg}
}
// C library error code constants - mapped from native.h
const (
// Special status codes
cFunctionOK = 0
cSendRequestParamOK = -15002 // Key not found, request sent via goProxy
// Init errors (-20000 ~ -20016)
cInitParamPathNull = -20000
cInitParamMyIDNull = -20001
cInitPathUTF8Null = -20002
cInitPathNotAvailable = -20003
cInitLogInitError = -20004
cInitEncryptInitError = -20005
cInitPCNativeError = -20006
cInitReadLocalKey = -20013
cInitReadKey = -20014
// EncryptData errors (-20100 ~ -20114)
cEncryptDataCorpIDNull = -20100
cEncryptDataMsgNull = -20101
cEncryptDataBufNull = -20102
cEncryptDataRetLenNull = -20103
cEncryptDataProxyNull = -20104
cEncryptDataLenError = -20105
cEncryptDataKeyNegtive = -20111
cEncryptDataTmpKeyNull = -20109
cEncryptDataBuildReqErr = -20110
// DecryptData errors (-20200 ~ -20215)
cDecryptDataCorpIDNull = -20200
cDecryptDataMsgNull = -20201
cDecryptDataBufNull = -20202
cDecryptDataRetLenNull = -20203
cDecryptDataProxyNull = -20204
cDecryptDataLenError = -20205
cDecryptDataFormatError = -20206
cDecryptDataDecryptErr = -20211
cDecryptDataKeyNegtive = -20214
cDecryptDataTmpKeyNull = -20212
cDecryptDataBuildReqErr = -20213
// EncryptFile errors (-20300 ~ -20339)
cEncryptFileCorpIDNull = -20300
cEncryptFilePathNull = -20301
cEncryptFileTmpKeyNull = -20315
cEncryptFileBuildReqErr = -20316
// EncryptBuffer errors (-20400 ~ -20416)
cEncryptBufferCorpIDNull = -20400
cEncryptBufferTmpKeyNull = -20414
cEncryptBufferBuildReqErr = -20415
// DecryptFile errors (-20500 ~ -20538)
cDecryptFileCorpIDNull = -20500
cDecryptFilePathNull = -20501
cDecryptFileHeadError = -20505
cDecryptFileHashError = -20517
cDecryptFileTmpKeyNull = -20518
cDecryptFileBuildReqErr = -20519
// DecryptBuffer errors (-20600 ~ -20622)
cDecryptBufferCorpIDNull = -20600
cDecryptBufferHeadNull = -20608
cDecryptBufferHashError = -20618
cDecryptBufferTmpKeyNull = -20619
cDecryptBufferBuildReqErr = -20620
// SetResponse errors (-20700 ~ -20730)
cSetResponseCorpIDNull = -20700
cSetResponseJSONNull = -20701
cSetResponseKeyNagtive = -20728
cSetResponseSaveKeyErr = -20723
// SetPushData errors (-20800 ~ -20814)
cSetPushDataCorpIDNull = -20800
cSetPushDataTypeUndef = -20814
// V3 signing errors (-31001 ~ -34002)
cV3BuildReq3SM2Failed = -31006
cV3BuildReq3SignError = -31009
cV3ParseResp3Downgrade = -32012
)
// cErrorMessages maps C error codes to human-readable messages.
var cErrorMessages = map[int]string{
// Init
-20000: "init: path parameter is NULL",
-20001: "init: my_id parameter is NULL",
-20002: "init: path UTF8 conversion returned NULL",
-20003: "init: path is not available/accessible",
-20004: "init: log initialization failed",
-20005: "init: encryption engine initialization failed",
-20006: "init: pcnative initialization failed",
-20007: "init: get URL full path error",
-20013: "init: read local key failed",
-20014: "init: read key failed",
-20015: "init: my_id parameter is NULL",
-20016: "init: logFunc parameter is NULL",
// EncryptData
-20100: "encryptData: corp_id is NULL or empty",
-20101: "encryptData: message content is NULL",
-20102: "encryptData: encrypt_buf is NULL",
-20103: "encryptData: ret_len is NULL",
-20104: "encryptData: proxy function is NULL",
-20105: "encryptData: data length invalid (too small or too big)",
-20106: "encryptData: encryptMsgHelper parameter error",
-20107: "encryptData: SM4 encryption failed",
-20108: "encryptData: base64 encoding failed",
-20109: "encryptData: key request already in progress",
-20110: "encryptData: build key request failed",
-20111: "encryptData: enterprise key is restricted",
-20113: "encryptData: malloc encode buffer failed",
-20114: "encryptData: malloc encrypt buffer failed",
// DecryptData
-20200: "decryptData: corp_id is NULL or empty",
-20201: "decryptData: message content is NULL",
-20202: "decryptData: decrypt_buf is NULL",
-20203: "decryptData: ret_len is NULL",
-20204: "decryptData: proxy function is NULL",
-20205: "decryptData: data length invalid",
-20206: "decryptData: message content format error (missing ||separators)",
-20207: "decryptData: decryptMsgHelper parameter error",
-20208: "decryptData: base64 decode failed",
-20209: "decryptData: decode length error",
-20210: "decryptData: decrypt buffer malloc failed",
-20211: "decryptData: SM4 decryption failed",
-20212: "decryptData: key request already in progress",
-20213: "decryptData: build key request failed",
-20214: "decryptData: enterprise key is restricted",
// EncryptFile
-20300: "encryptFile: corp_id is NULL or empty",
-20301: "encryptFile: source or dest file path is NULL",
-20302: "encryptFile: id parameter is NULL",
-20303: "encryptFile: seq_id parameter is NULL",
-20304: "encryptFile: proxy function is NULL",
-20305: "encryptFile: key_info is NULL",
-20308: "encryptFile: get file size failed",
-20309: "encryptFile: open source or dest file failed",
-20313: "encryptFile: encrypt block failed",
-20315: "encryptFile: key request already in progress",
-20316: "encryptFile: build key request failed",
-20333: "encryptFile: create thread failed",
-20339: "encryptFile: test environment error",
// EncryptBuffer
-20400: "encryptBuffer: corp_id is NULL or empty",
-20401: "encryptBuffer: file content is NULL",
-20407: "encryptBuffer: proxy function is NULL",
-20413: "encryptBuffer: encrypted block error",
-20414: "encryptBuffer: key request already in progress",
-20415: "encryptBuffer: build key request failed",
// DecryptFile
-20500: "decryptFile: corp_id is NULL or empty",
-20501: "decryptFile: source or dest file path is NULL",
-20505: "decryptFile: file header format error",
-20509: "decryptFile: header magic error",
-20511: "decryptFile: open file failed",
-20515: "decryptFile: decrypt block failed",
-20517: "decryptFile: file hash verification failed (warning)",
-20518: "decryptFile: key request already in progress",
-20519: "decryptFile: build key request failed",
-20520: "decryptFile: file length error",
// DecryptBuffer
-20600: "decryptBuffer: corp_id is NULL or empty",
-20608: "decryptBuffer: parseHeadInfo returned NULL",
-20615: "decryptBuffer: header magic error",
-20617: "decryptBuffer: decrypt block error",
-20618: "decryptBuffer: file hash verification failed",
-20619: "decryptBuffer: key request already in progress",
-20620: "decryptBuffer: build key request failed",
-20621: "decryptBuffer: file length error",
// SetResponse
-20700: "setResponse: corp_id is NULL or empty",
-20701: "setResponse: json_str is NULL or empty",
-20703: "setResponse: find tmp corp no tmp key found",
-20704: "setResponse: JSON parse failed",
-20716: "setResponse: server request error",
-20722: "setResponse: localKeyPath is NULL",
-20723: "setResponse: save key failed",
-20728: "setResponse: enterprise key is restricted (nagtive)",
// SetPushData
-20800: "setPushData: corp_id is NULL",
-20801: "setPushData: push_data is NULL",
-20804: "setPushData: JSON parse failed",
-20805: "setPushData: type field is NULL",
-20806: "setPushData: find_create_tmp_key failed",
-20814: "setPushData: unknown push type",
// V3 Signing
-31001: "v3: build_key_request3 arg is NULL",
-31006: "v3: SM2 encrypt R2 failed",
-31009: "v3: sign step one error",
-32005: "v3: deserialize response failed",
-32010: "v3: new SDK request old private error",
-32012: "v3: public server downgrade",
-32022: "v3: MAC of response mismatch",
-32029: "v3: verify server sign failed",
// Misc
-34001: "v3: setResponse not found tmp sign",
-34002: "v3: setResponse malloc for ret buf failed",
}
+463
View File
@@ -0,0 +1,463 @@
// Package main demonstrates how to integrate the SafeChat Go SDK
// into the DingTalk Workspace CLI or any other Go application.
//
// Two usage modes:
//
// 1. Single-action mode (default)
// Run one of: encrypt-msg / decrypt-msg / encrypt-file / decrypt-file /
// encrypt-buf / decrypt-buf via -action flag.
//
// 2. Full-test mode (-test-all)
// Exercise all 3 encryption APIs (Msg / File / Buffer) in one run,
// verify round-trip (decrypt == original), and print a summary.
//
// Key cache files
//
// The SDK persists negotiated keys under -data directory:
// ahflag_256.store local key identifier
// ahkey_256.store encrypted key material
//
// Build:
//
// go build -o safechat-example ./example/
package main
import (
"bytes"
"crypto/rand"
"flag"
"fmt"
"log"
"os"
"path/filepath"
"strings"
"time"
safechat "safechat-go-sdk"
)
func main() {
var (
dataPath = flag.String("data", "./keystore", "Path to key storage directory (contains ahflag_256.store / ahkey_256.store)")
userID = flag.String("user", "", "Current login user ID")
code = flag.String("code", "", "DingTalk authCode (only needed when keys must be fetched from server)")
corpID = flag.String("corp", "", "Enterprise/corp ID (required)")
staffID = flag.String("staff", "", "Staff ID for encryption target (defaults to -user)")
action = flag.String("action", "", "Single action: encrypt-msg|decrypt-msg|encrypt-file|decrypt-file|encrypt-buf|decrypt-buf")
input = flag.String("input", "", "Input text / file path")
output = flag.String("output", "", "Output file path (for file operations)")
server = flag.String("server", "", "Key server URL override (optional)")
testAll = flag.Bool("test-all", false, "Run a full round-trip test of all 3 encryption APIs using cached keys")
verbose = flag.Bool("v", false, "Verbose logging")
)
flag.Parse()
if *corpID == "" {
fmt.Fprintln(os.Stderr, "Error: -corp flag is required")
flag.Usage()
os.Exit(1)
}
// Validate key cache files up-front so the user gets a clear message
// instead of a cryptic C-library error later.
checkKeyCache(*dataPath, *code)
cfg := safechat.Config{
DataPath: *dataPath,
UserID: *userID,
Code: *code,
MaxRetry: 5,
HTTPTimeout: 15 * time.Second,
}
if !*verbose {
cfg.Logger = &quietLogger{}
} else {
cfg.Logger = &stdLogger{}
}
if *server != "" {
cfg.KeyServer = *server
}
client, err := safechat.New(cfg)
if err != nil {
log.Fatalf("Failed to initialize SafeChat client: %v", err)
}
defer client.Close()
targetStaff := *staffID
if targetStaff == "" {
targetStaff = *userID
}
if *testAll {
runFullTest(client, *corpID, targetStaff)
return
}
if *action == "" {
fmt.Fprintln(os.Stderr, "Error: either -action <name> or -test-all must be provided")
flag.Usage()
os.Exit(1)
}
switch *action {
case "encrypt-msg":
encryptMessage(client, *corpID, targetStaff, *input)
case "decrypt-msg":
decryptMessage(client, *corpID, targetStaff, *input)
case "encrypt-file":
encryptFile(client, *corpID, targetStaff, *input, *output)
case "decrypt-file":
decryptFile(client, *corpID, targetStaff, *input, *output)
case "encrypt-buf":
encryptBuffer(client, *corpID, targetStaff, *input)
case "decrypt-buf":
decryptBuffer(client, *corpID, targetStaff, *input)
default:
fmt.Fprintf(os.Stderr, "Unknown action: %s\n", *action)
os.Exit(1)
}
}
// ---------- key-cache pre-check ----------
// checkKeyCache prints a friendly hint about which keys are available and
// whether a network round-trip to the key server should be expected.
func checkKeyCache(dataPath, code string) {
flagFile := filepath.Join(dataPath, "ahflag_256.store")
keyFile := filepath.Join(dataPath, "ahkey_256.store")
_, err1 := os.Stat(flagFile)
_, err2 := os.Stat(keyFile)
hasFlag := err1 == nil
hasKey := err2 == nil
switch {
case hasFlag && hasKey:
fmt.Printf("[key-cache] found %s + %s (existence check only; C layer still validates hash/version)\n",
filepath.Base(flagFile), filepath.Base(keyFile))
fmt.Printf("[key-cache] NOTE: store files are NOT portable across CPU architectures — " +
"calc_hash() depends on char signedness (signed on x86_64, unsigned on AArch64). " +
"A foreign store fails the hash check, gets deleted and re-generated, which triggers a server call.\n")
case hasFlag || hasKey:
fmt.Fprintf(os.Stderr,
"[key-cache] WARNING: only one of the pair exists (%s / %s); key negotiation will likely fail\n",
filepath.Base(flagFile), filepath.Base(keyFile))
default:
if code == "" {
fmt.Fprintf(os.Stderr,
"[key-cache] no cached keys in %s and -code is empty; "+
"either copy ahflag_256.store + ahkey_256.store into that dir, "+
"or provide a valid DingTalk access_token via -code\n",
dataPath)
} else {
fmt.Printf("[key-cache] no cached keys in %s; will try to fetch from server using -code\n", dataPath)
}
}
}
// ---------- full round-trip test of all 3 APIs ----------
// runFullTest exercises the 3 Go API pairs defined in safechat.go.
//
// The APIs under test are the PUBLIC Go methods on *safechat.Client:
//
// API #1 — Msg API : EncryptMsg / DecryptMsg (safechat.go L115, L154)
// API #2 — File API : EncryptFile / DecryptFile (safechat.go L188, L213)
// API #3 — Buffer API : EncryptBuffer / DecryptBuffer (safechat.go L240, L268)
//
// These are pure Go entry points: parameter validation, mutex locking,
// MaxRetry loop, and Go error translation are all done in safechat.go.
// The underlying CGO bridge (cEncryptData / cDecryptData / ...) is a
// PRIVATE implementation detail and is NOT what this test targets.
func runFullTest(client *safechat.Client, corpID, staffID string) {
fmt.Println("==========================================================")
fmt.Printf(" SafeChat Go SDK — Go API round-trip test\n")
fmt.Printf(" corpID=%s staffID=%s\n", corpID, staffID)
fmt.Println(" Target: the 3 public Go API pairs in safechat.go")
fmt.Println(" #1 Msg API : EncryptMsg / DecryptMsg")
fmt.Println(" #2 File API : EncryptFile / DecryptFile")
fmt.Println(" #3 Buffer API : EncryptBuffer / DecryptBuffer")
fmt.Println("==========================================================")
passed, failed := 0, 0
// Go API #1 — Msg
fmt.Println("\n[1/3] Go API #1 : EncryptMsg / DecryptMsg (safechat.go L115, L154)")
if runCase("Msg API", func() error { return testMsgAPI(client, corpID, staffID) }) {
passed++
} else {
failed++
}
// Go API #2 — File
fmt.Println("\n[2/3] Go API #2 : EncryptFile / DecryptFile (safechat.go L188, L213)")
if runCase("File API", func() error { return testFileAPI(client, corpID, staffID) }) {
passed++
} else {
failed++
}
// Go API #3 — Buffer
fmt.Println("\n[3/3] Go API #3 : EncryptBuffer / DecryptBuffer (safechat.go L240, L268)")
if runCase("Buffer API", func() error { return testBufferAPI(client, corpID, staffID) }) {
passed++
} else {
failed++
}
fmt.Println("\n==========================================================")
fmt.Printf(" RESULT: %d passed, %d failed\n", passed, failed)
fmt.Println("==========================================================")
if failed > 0 {
os.Exit(1)
}
}
func runCase(name string, fn func() error) bool {
if err := fn(); err != nil {
fmt.Printf(" ✗ %s FAILED: %v\n", name, err)
return false
}
fmt.Printf(" ✓ %s OK\n", name)
return true
}
// testMsgAPI exercises Go API #1:
//
// client.EncryptMsg(corpID, staffID, plain) → ciphertext
// client.DecryptMsg(corpID, staffID, ciphertext) → plaintext
//
// Both methods are defined in safechat.go (L115, L154). They perform
// parameter validation, mutex locking, and a MaxRetry loop before
// returning a Go []byte + error.
func testMsgAPI(c *safechat.Client, corpID, staffID string) error {
plain := make([]byte, 128)
if _, err := rand.Read(plain); err != nil {
return err
}
// Go API call — safechat.go L115
ct, err := c.EncryptMsg(corpID, staffID, plain)
if err != nil {
return fmt.Errorf("EncryptMsg (Go API): %w", err)
}
fmt.Printf(" [Go API] EncryptMsg : plain=%d bytes -> ct=%d bytes\n", len(plain), len(ct))
fmt.Printf(" [ciphertext str] %s\n", strings.ReplaceAll(string(ct), "\n", ""))
// Go API call — safechat.go L154
got, err := c.DecryptMsg(corpID, staffID, ct)
if err != nil {
return fmt.Errorf("DecryptMsg (Go API): %w", err)
}
if !bytes.Equal(got, plain) {
return fmt.Errorf("round-trip mismatch: got %d bytes, want %d", len(got), len(plain))
}
fmt.Printf(" [Go API] DecryptMsg : ct=%d bytes -> plain=%d bytes OK\n", len(ct), len(got))
return nil
}
// testFileAPI exercises Go API #2:
//
// client.EncryptFile(corpID, staffID, src, enc) error
// client.DecryptFile(corpID, staffID, enc, dec) error
//
// Both methods are defined in safechat.go (L188, L213). They operate on
// file paths and return only an error — no data crosses the Go/C boundary
// in the caller-visible API.
func testFileAPI(c *safechat.Client, corpID, staffID string) error {
dir, err := os.MkdirTemp("", "safechat-file-test-*")
if err != nil {
return err
}
defer os.RemoveAll(dir)
src := filepath.Join(dir, "plain.bin")
enc := filepath.Join(dir, "plain.bin.enc")
dec := filepath.Join(dir, "plain.bin.dec")
buf := make([]byte, 1024)
if _, err := rand.Read(buf); err != nil {
return err
}
if err := os.WriteFile(src, buf, 0644); err != nil {
return err
}
// Go API call — safechat.go L188
if err := c.EncryptFile(corpID, staffID, src, enc); err != nil {
return fmt.Errorf("EncryptFile (Go API): %w", err)
}
encSize, _ := fileSize(enc)
fmt.Printf(" [Go API] EncryptFile : src=%d bytes -> enc=%d bytes\n", len(buf), encSize)
// Go API call — safechat.go L213
if err := c.DecryptFile(corpID, staffID, enc, dec); err != nil {
return fmt.Errorf("DecryptFile (Go API): %w", err)
}
decBuf, err := os.ReadFile(dec)
if err != nil {
return err
}
if !bytes.Equal(decBuf, buf) {
return fmt.Errorf("round-trip mismatch: dec=%d bytes, want=%d", len(decBuf), len(buf))
}
fmt.Printf(" [Go API] DecryptFile : enc=%d bytes -> dec=%d bytes OK\n", encSize, len(decBuf))
return nil
}
// testBufferAPI exercises Go API #3:
//
// client.EncryptBuffer(corpID, staffID, data) → []byte, error
// client.DecryptBuffer(corpID, staffID, data) → []byte, error
//
// Both methods are defined in safechat.go (L240, L268). They operate on
// in-memory []byte buffers, suitable for binary protocols or DB blobs.
func testBufferAPI(c *safechat.Client, corpID, staffID string) error {
plain := make([]byte, 512)
if _, err := rand.Read(plain); err != nil {
return err
}
// Go API call — safechat.go L240
ct, err := c.EncryptBuffer(corpID, staffID, plain)
if err != nil {
return fmt.Errorf("EncryptBuffer (Go API): %w", err)
}
fmt.Printf(" [Go API] EncryptBuffer : plain=%d bytes -> ct=%d bytes\n", len(plain), len(ct))
// Go API call — safechat.go L268
got, err := c.DecryptBuffer(corpID, staffID, ct)
if err != nil {
return fmt.Errorf("DecryptBuffer (Go API): %w", err)
}
if !bytes.Equal(got, plain) {
return fmt.Errorf("round-trip mismatch: got %d bytes, want %d", len(got), len(plain))
}
fmt.Printf(" [Go API] DecryptBuffer : ct=%d bytes -> plain=%d bytes OK\n", len(ct), len(got))
return nil
}
func fileSize(p string) (int64, error) {
fi, err := os.Stat(p)
if err != nil {
return 0, err
}
return fi.Size(), nil
}
// ---------- single-action helpers ----------
func encryptMessage(client *safechat.Client, corpID, staffID, plaintext string) {
if plaintext == "" {
plaintext = "Hello, this is a test message from SafeChat Go SDK!"
}
fmt.Printf("Encrypting message: %q\n", plaintext)
ciphertext, err := client.EncryptMsg(corpID, staffID, []byte(plaintext))
if err != nil {
log.Fatalf("EncryptMsg failed: %v", err)
}
fmt.Printf("Encrypted (%d bytes): %s\n", len(ciphertext), string(ciphertext))
}
func decryptMessage(client *safechat.Client, corpID, staffID, ciphertext string) {
if ciphertext == "" {
log.Fatal("decrypt-msg requires -input with the ciphertext")
}
fmt.Printf("Decrypting message (%d bytes)...\n", len(ciphertext))
plaintext, err := client.DecryptMsg(corpID, staffID, []byte(ciphertext))
if err != nil {
log.Fatalf("DecryptMsg failed: %v", err)
}
fmt.Printf("Decrypted: %s\n", string(plaintext))
}
func encryptFile(client *safechat.Client, corpID, staffID, srcPath, dstPath string) {
if srcPath == "" {
log.Fatal("encrypt-file requires -input with source file path")
}
if dstPath == "" {
dstPath = srcPath + ".enc"
}
fmt.Printf("Encrypting file: %s -> %s\n", srcPath, dstPath)
if err := client.EncryptFile(corpID, staffID, srcPath, dstPath); err != nil {
log.Fatalf("EncryptFile failed: %v", err)
}
fmt.Println("File encrypted successfully")
}
func decryptFile(client *safechat.Client, corpID, staffID, srcPath, dstPath string) {
if srcPath == "" {
log.Fatal("decrypt-file requires -input with source file path")
}
if dstPath == "" {
dstPath = srcPath + ".dec"
}
fmt.Printf("Decrypting file: %s -> %s\n", srcPath, dstPath)
if err := client.DecryptFile(corpID, staffID, srcPath, dstPath); err != nil {
log.Fatalf("DecryptFile failed: %v", err)
}
fmt.Println("File decrypted successfully")
}
func encryptBuffer(client *safechat.Client, corpID, staffID, inputPath string) {
if inputPath == "" {
log.Fatal("encrypt-buf requires -input with file path containing data to encrypt")
}
data, err := os.ReadFile(inputPath)
if err != nil {
log.Fatalf("Failed to read input file: %v", err)
}
fmt.Printf("Encrypting buffer (%d bytes)...\n", len(data))
encrypted, err := client.EncryptBuffer(corpID, staffID, data)
if err != nil {
log.Fatalf("EncryptBuffer failed: %v", err)
}
outPath := inputPath + ".enc"
if err := os.WriteFile(outPath, encrypted, 0644); err != nil {
log.Fatalf("Failed to write output: %v", err)
}
fmt.Printf("Buffer encrypted successfully (%d bytes) -> %s\n", len(encrypted), outPath)
}
func decryptBuffer(client *safechat.Client, corpID, staffID, inputPath string) {
if inputPath == "" {
log.Fatal("decrypt-buf requires -input with file path containing data to decrypt")
}
data, err := os.ReadFile(inputPath)
if err != nil {
log.Fatalf("Failed to read input file: %v", err)
}
fmt.Printf("Decrypting buffer (%d bytes)...\n", len(data))
decrypted, err := client.DecryptBuffer(corpID, staffID, data)
if err != nil {
log.Fatalf("DecryptBuffer failed: %v", err)
}
outPath := inputPath + ".dec"
if err := os.WriteFile(outPath, decrypted, 0644); err != nil {
log.Fatalf("Failed to write output: %v", err)
}
fmt.Printf("Buffer decrypted successfully (%d bytes) -> %s\n", len(decrypted), outPath)
}
// ---------- loggers ----------
type stdLogger struct{}
func (l *stdLogger) Debug(format string, args ...interface{}) {
log.Printf("[DEBUG] "+format, args...)
}
func (l *stdLogger) Info(format string, args ...interface{}) {
log.Printf("[INFO] "+format, args...)
}
func (l *stdLogger) Error(format string, args ...interface{}) {
log.Printf("[ERROR] "+format, args...)
}
// quietLogger swallows SDK logs; only errors are surfaced via returned error values.
type quietLogger struct{}
func (l *quietLogger) Debug(format string, args ...interface{}) {}
func (l *quietLogger) Info(format string, args ...interface{}) {}
func (l *quietLogger) Error(format string, args ...interface{}) {}
+5
View File
@@ -0,0 +1,5 @@
module safechat-go-sdk
go 1.21
require github.com/google/uuid v1.6.0
+2
View File
@@ -0,0 +1,2 @@
github.com/google/uuid v1.6.0 h1:NIvaJDMOsjHA8n1jAhLSgzrAzy1Hgr+hNrb57e+94F0=
github.com/google/uuid v1.6.0/go.mod h1:TIyPZe4MgqvfeYDBFedMoGGpEw/LqOeaOT+nhxU+yHo=
+80
View File
@@ -0,0 +1,80 @@
#include "goproxy_bridge.h"
#include "csrc/safechat.h"
#include <stdlib.h>
/*
* goproxy_bridge.c - CGO callback bridge implementation
*
* This file MUST reside in the Go package root directory so that
* CGO compiles it together with the Go code. It bridges C library
* callback invocations to Go exported functions.
*
* Go exported functions are declared as extern here.
* The bridge functions simply forward calls from C library
* to the Go runtime via CGO mechanism.
*/
/* Go exported function declarations (defined in callback.go) */
extern int goProxy(char *corpid, char *uid, char *domain,
char *url, char *param, char *seq_id);
extern int goBlock(char *corpid);
extern int goCancelBlock(char *corpid);
/*
* init - Wrapper for safechatInit
* Initializes the SafeChat library.
*/
int init(char *path, char *my_id, void *reserved) {
(void)reserved; /* unused */
return safechatInit(path, my_id);
}
/*
* clearCache - Clears the key cache for a specific enterprise
* Returns 0 on success, non-zero on error.
* Note: This is a placeholder - actual implementation may vary.
*/
int clearCache(char *corpid) {
(void)corpid; /* unused for now */
/* TODO: Implement actual cache clearing if needed */
return 0;
}
/*
* freeCryptoBuf - Frees a buffer allocated by the C library
* The C library allocates buffers with malloc, so we free with free.
*/
void freeCryptoBuf(void *buf) {
if (buf != NULL) {
free(buf);
}
}
/*
* goProxyBridge - Bridge for call_proxy_func typedef
* Called by C library when a key request needs to be sent.
* Forwards to Go's goProxy which performs HTTP request
* and calls setResponse to feed back the key data.
*/
int goProxyBridge(char *corpid, char *uid, char *domain,
char *url, char *param, char *seq_id) {
return goProxy(corpid, uid, domain, url, param, seq_id);
}
/*
* goBlockBridge - Bridge for block_crypto_func typedef
* Called by C library when enterprise key is restricted.
* Notifies Go layer of the restriction.
*/
int goBlockBridge(char *corpid) {
return goBlock(corpid);
}
/*
* goCancelBlockBridge - Bridge for cancel_block_crypto_func typedef
* Called by C library when enterprise key restriction is lifted.
* Notifies Go layer to remove the restriction.
*/
int goCancelBlockBridge(char *corpid) {
return goCancelBlock(corpid);
}
+28
View File
@@ -0,0 +1,28 @@
#ifndef GOPROXY_BRIDGE_H
#define GOPROXY_BRIDGE_H
/*
* goproxy_bridge.h - CGO callback bridge declarations
*
* These bridge functions are called by the C library (safechat.c)
* and forward to Go exported functions via CGO.
* This indirection is required because CGO cannot directly pass
* Go function pointers to C code.
*/
/* Proxy callback bridge - forwards key requests to Go HTTP client */
int goProxyBridge(char *corpid, char *uid, char *domain,
char *url, char *param, char *seq_id);
/* Block crypto callback bridge - notifies Go of key restriction */
int goBlockBridge(char *corpid);
/* Cancel block crypto callback bridge - notifies Go of restriction lift */
int goCancelBlockBridge(char *corpid);
/* Wrapper functions for CGO */
int init(char *path, char *my_id, void *reserved);
int clearCache(char *corpid);
void freeCryptoBuf(void *buf);
#endif /* GOPROXY_BRIDGE_H */
+155
View File
@@ -0,0 +1,155 @@
package safechat
import (
"crypto/tls"
"fmt"
"io"
"net/http"
"net/url"
"strings"
"time"
)
// httpLogBodyLimit is the maximum body length (bytes) included in debug logs.
// Beyond this, the body is truncated with a "..." marker so the log line stays
// readable. Full bodies are still passed to the C library unchanged.
const httpLogBodyLimit = 4096
// keyClient handles HTTP communication with the key server.
// It is responsible for sending key requests (triggered by C library's goProxy callback)
// and returning the JSON response to be fed back via setResponse.
//
// The code field holds a DingTalk authCode used for server authentication.
// It is required when:
// - Fetching keys for the first time (empty keystore)
// - Server-side key version rotation (C library detects version mismatch)
type keyClient struct {
httpClient *http.Client
code string // DingTalk authCode
keyServer string // Optional override for key server URL
logger Logger
}
// newKeyClient creates a new key client with the given configuration.
func newKeyClient(cfg Config) *keyClient {
transport := &http.Transport{
TLSClientConfig: &tls.Config{
// Allow connecting to enterprise private servers with self-signed certs
InsecureSkipVerify: true,
},
MaxIdleConns: 10,
IdleConnTimeout: 30 * time.Second,
DisableCompression: true,
}
return &keyClient{
httpClient: &http.Client{
Timeout: cfg.HTTPTimeout,
Transport: transport,
},
code: cfg.Code,
keyServer: cfg.KeyServer,
logger: cfg.Logger,
}
}
func (kc *keyClient) doKeyRequest(fullURL, param, code string) (string, error) {
// Use override key server if configured
targetURL := fullURL
if kc.keyServer != "" {
targetURL = kc.keyServer
}
var body string
if strings.HasPrefix(param, "param=") {
// V1 format: C library output already has "param=" prefix and all fields.
body = fmt.Sprintf("%s&code=%s", param, url.QueryEscape(code))
} else {
body = fmt.Sprintf("%s&code=%s&appAlgVersion=1", param, url.QueryEscape(code))
}
// === Request logging (debug) ===
// Log full request line, headers and body (truncated) so we can verify
// the URL, the URL-encoded payload and the auth code are shaped as expected.
if kc.logger != nil {
kc.logger.Debug("=== HTTP key request ===")
kc.logger.Debug("URL: POST %s", targetURL)
kc.logger.Debug("C-URL: %s (from C library)", fullURL)
kc.logger.Debug("Server: %s%s",
func() string {
if kc.keyServer != "" {
return kc.keyServer + " (override)"
}
return "(C-provided)"
}(),
"")
kc.logger.Debug("Headers: Content-Type=application/x-www-form-urlencoded, User-Agent=SafeChat-Go-SDK/1.0")
kc.logger.Debug("Code (auth_token, length=%d): %s", len(code), previewString(code, 64))
kc.logger.Debug("Body (length=%d): %s", len(body), previewString(body, httpLogBodyLimit))
}
req, err := http.NewRequest("POST", targetURL, strings.NewReader(body))
if err != nil {
if kc.logger != nil {
kc.logger.Error("create request failed: %v", err)
}
return "", fmt.Errorf("create request failed: %w", err)
}
req.Header.Set("Content-Type", "application/x-www-form-urlencoded")
req.Header.Set("User-Agent", "SafeChat-Go-SDK/1.0")
// Capture timing so we can see if the server is slow / timing out.
start := time.Now()
resp, err := kc.httpClient.Do(req)
if err != nil {
if kc.logger != nil {
kc.logger.Error("HTTP request failed after %s: %v", time.Since(start), err)
}
return "", fmt.Errorf("HTTP request failed: %w", err)
}
defer resp.Body.Close()
respBody, err := io.ReadAll(resp.Body)
if err != nil {
if kc.logger != nil {
kc.logger.Error("read response body failed after %s: %v", time.Since(start), err)
}
return "", fmt.Errorf("read response body failed: %w", err)
}
duration := time.Since(start)
// === Response logging (debug on success, info on non-2xx) ===
if kc.logger != nil {
kc.logger.Debug("=== HTTP key response ===")
kc.logger.Debug("Status: %d %s (took %s)", resp.StatusCode, http.StatusText(resp.StatusCode), duration)
kc.logger.Debug("Headers: Content-Type=%s, Content-Length=%d", resp.Header.Get("Content-Type"), len(respBody))
kc.logger.Debug("Body (length=%d): %s", len(respBody), previewString(string(respBody), httpLogBodyLimit))
}
if resp.StatusCode != http.StatusOK {
if kc.logger != nil {
kc.logger.Error("key server returned HTTP %d %s (took %s, body length=%d): %s",
resp.StatusCode, http.StatusText(resp.StatusCode), duration, len(respBody),
previewString(string(respBody), httpLogBodyLimit))
}
return "", fmt.Errorf("key server returned HTTP %d", resp.StatusCode)
}
return string(respBody), nil
}
// previewString returns s unchanged when it fits in max bytes; otherwise it
// returns the first max bytes followed by "...[truncated, total=N]". This is
// used to keep HTTP request/response log lines readable for very large bodies
// while still preserving the head of the payload for debugging.
func previewString(s string, max int) string {
if max <= 0 || len(s) <= max {
return s
}
return s[:max] + fmt.Sprintf("...[truncated, total=%d]", len(s))
}
// updateCode updates the authentication code (may change during runtime).
func (kc *keyClient) updateCode(code string) {
kc.code = code
}
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
+615
View File
@@ -0,0 +1,615 @@
// Package safechat provides encryption/decryption capabilities for the
// Basic usage:
//
// client, err := safechat.New(safechat.Config{
// DataPath: "/path/to/keystore",
// UserID: "user123",
// Code: "authCode ", // authCode
// })
// if err != nil {
// log.Fatal(err)
// }
// defer client.Close()
//
// ciphertext, err := client.EncryptMsg("corp_id", "staff_id", []byte("hello"))
//
// The Code field must be a valid DingTalk authCode. It is required
// when fetching keys from the server (first-time use or key version rotation).
package safechat
/*
#include "csrc/safechat.h"
#include "goproxy_bridge.h"
#include <stdlib.h>
#include <string.h>
*/
import "C"
import (
"fmt"
"sync"
"unsafe"
"github.com/google/uuid"
)
// Client is the main SafeChat encryption client.
// It wraps the C library and provides a thread-safe Go API.
//
// IMPORTANT: Only one Client instance can exist per process because
// the underlying C library uses global state. Creating a second Client
// will return ErrAlreadyInitialized.
//
// All public methods are safe for concurrent use from multiple goroutines.
// Internal synchronization uses a two-lock design:
// - mu: serializes all C library calls (prevents C global state corruption)
// - kcMu: protects keyClient state during HTTP requests (used in goProxy callback)
type Client struct {
cfg Config
mu sync.Mutex // Serializes all C library calls
kcMu sync.Mutex // Protects keyClient during HTTP key requests
kc *keyClient // HTTP client for key server communication
inited bool // Whether C library init() has been called
blockedCorps sync.Map // map[string]bool - enterprises with restricted keys
}
// New creates a new SafeChat client and initializes the underlying C library.
//
// The Config.DataPath directory will be used to store encryption keys and
// related metadata. It must exist and be writable.
//
// Returns ErrAlreadyInitialized if called more than once per process.
func New(cfg Config) (*Client, error) {
cfg = defaultConfig(cfg)
if err := cfg.validate(); err != nil {
return nil, err
}
// Check if already initialized (C library is singleton)
if getGlobalClient() != nil {
return nil, ErrAlreadyInitialized
}
c := &Client{
cfg: cfg,
kc: newKeyClient(cfg),
}
// Register globally for CGO callbacks before calling init
registerGlobalClient(c)
// Initialize C library
if err := c.cInit(); err != nil {
globalClient.Store((*Client)(nil))
return nil, fmt.Errorf("safechat init failed: %w", err)
}
c.inited = true
return c, nil
}
// Close releases resources held by the client.
// After Close is called, no other methods should be called.
func (c *Client) Close() {
c.mu.Lock()
defer c.mu.Unlock()
c.inited = false
// Clear the global singleton so a new Client can be created after Close.
// The C library has no explicit cleanup function; keys are persisted to
// disk and process memory is reclaimed on exit.
globalClient.Store((*Client)(nil))
}
// EncryptMsg encrypts a plaintext message for the given enterprise.
//
// Returns the ciphertext in the standard SafeChat format:
// base64(encrypted_data)||key_version||method_num||plain_length
//
// If the encryption key is not yet available, the SDK will automatically
// request it from the key server (via goProxy callback) and retry.
func (c *Client) EncryptMsg(corpID, staffID string, plaintext []byte) ([]byte, error) {
if !c.inited {
return nil, ErrNotInitialized
}
if len(plaintext) == 0 {
return nil, fmt.Errorf("safechat: plaintext cannot be empty")
}
c.mu.Lock()
defer c.mu.Unlock()
for i := 0; i <= c.cfg.MaxRetry; i++ {
result, err := c.cEncryptData(corpID, staffID, plaintext)
if err == nil {
return result, nil
}
// Check if it's a "key requested" status - retry
if cerr, ok := err.(*CError); ok && cerr.Code == cSendRequestParamOK {
// goProxy was called and setResponse was invoked synchronously.
// Next iteration should find the key in local cache.
continue
}
// Check for key restriction
if cerr, ok := err.(*CError); ok && cerr.Code == cEncryptDataKeyNegtive {
return nil, ErrKeyRestricted
}
return nil, err
}
return nil, ErrMaxRetryExceeded
}
// DecryptMsg decrypts a ciphertext message.
//
// The ciphertext must be in the standard SafeChat format:
// base64(encrypted_data)||key_version||method_num||plain_length
func (c *Client) DecryptMsg(corpID, staffID string, ciphertext []byte) ([]byte, error) {
if !c.inited {
return nil, ErrNotInitialized
}
if len(ciphertext) == 0 {
return nil, fmt.Errorf("safechat: ciphertext cannot be empty")
}
c.mu.Lock()
defer c.mu.Unlock()
for i := 0; i <= c.cfg.MaxRetry; i++ {
result, err := c.cDecryptData(corpID, staffID, ciphertext)
if err == nil {
return result, nil
}
if cerr, ok := err.(*CError); ok && cerr.Code == cSendRequestParamOK {
continue
}
if cerr, ok := err.(*CError); ok && cerr.Code == cDecryptDataKeyNegtive {
return nil, ErrKeyRestricted
}
return nil, err
}
return nil, ErrMaxRetryExceeded
}
// EncryptFile encrypts a file from srcPath to dstPath.
//
// The encrypted file uses a 12-byte header (msg_HandInfo_t) followed by
// SM4-ECB encrypted data in 8KiB blocks.
func (c *Client) EncryptFile(corpID, staffID, srcPath, dstPath string) error {
if !c.inited {
return ErrNotInitialized
}
c.mu.Lock()
defer c.mu.Unlock()
for i := 0; i <= c.cfg.MaxRetry; i++ {
err := c.cEncryptFile(corpID, staffID, srcPath, dstPath)
if err == nil {
return nil
}
if cerr, ok := err.(*CError); ok && cerr.Code == cSendRequestParamOK {
continue
}
return err
}
return ErrMaxRetryExceeded
}
// DecryptFile decrypts a file from srcPath to dstPath.
func (c *Client) DecryptFile(corpID, staffID, srcPath, dstPath string) error {
if !c.inited {
return ErrNotInitialized
}
c.mu.Lock()
defer c.mu.Unlock()
for i := 0; i <= c.cfg.MaxRetry; i++ {
err := c.cDecryptFile(corpID, staffID, srcPath, dstPath)
if err == nil {
return nil
}
if cerr, ok := err.(*CError); ok && cerr.Code == cSendRequestParamOK {
continue
}
return err
}
return ErrMaxRetryExceeded
}
// EncryptBuffer encrypts binary data in memory.
//
// The result includes a 12-byte header followed by SM4-ECB encrypted content.
func (c *Client) EncryptBuffer(corpID, staffID string, data []byte) ([]byte, error) {
if !c.inited {
return nil, ErrNotInitialized
}
if len(data) == 0 {
return nil, fmt.Errorf("safechat: data cannot be empty")
}
c.mu.Lock()
defer c.mu.Unlock()
for i := 0; i <= c.cfg.MaxRetry; i++ {
result, err := c.cEncryptBuffer(corpID, staffID, data)
if err == nil {
return result, nil
}
if cerr, ok := err.(*CError); ok && cerr.Code == cSendRequestParamOK {
continue
}
return nil, err
}
return nil, ErrMaxRetryExceeded
}
// DecryptBuffer decrypts binary data in memory.
func (c *Client) DecryptBuffer(corpID, staffID string, data []byte) ([]byte, error) {
if !c.inited {
return nil, ErrNotInitialized
}
if len(data) == 0 {
return nil, fmt.Errorf("safechat: data cannot be empty")
}
c.mu.Lock()
defer c.mu.Unlock()
for i := 0; i <= c.cfg.MaxRetry; i++ {
result, err := c.cDecryptBuffer(corpID, staffID, data)
if err == nil {
return result, nil
}
if cerr, ok := err.(*CError); ok && cerr.Code == cSendRequestParamOK {
continue
}
return nil, err
}
return nil, ErrMaxRetryExceeded
}
// SetResponse manually injects a key server response into the C library.
// This is an advanced API for cases where the caller handles HTTP
// communication externally.
func (c *Client) SetResponse(corpID, jsonStr string) error {
if !c.inited {
return ErrNotInitialized
}
c.mu.Lock()
defer c.mu.Unlock()
return c.cSetResponse(corpID, jsonStr)
}
// SetPushData processes a server push notification (key update, etc).
func (c *Client) SetPushData(corpID, staffID, pushData string) error {
if !c.inited {
return ErrNotInitialized
}
c.mu.Lock()
defer c.mu.Unlock()
return c.cSetPushData(corpID, staffID, pushData)
}
// ClearCache clears the local key cache for a specific enterprise.
func (c *Client) ClearCache(corpID string) error {
if !c.inited {
return ErrNotInitialized
}
c.mu.Lock()
defer c.mu.Unlock()
cCorpID := C.CString(corpID)
defer C.free(unsafe.Pointer(cCorpID))
ret := C.clearCache(cCorpID)
return mapCError(int(ret))
}
// IsBlocked returns true if the given enterprise's key is restricted.
func (c *Client) IsBlocked(corpID string) bool {
_, ok := c.blockedCorps.Load(corpID)
return ok
}
// UpdateCode updates the DingTalk authentication code at runtime.
func (c *Client) UpdateCode(code string) {
c.kcMu.Lock()
defer c.kcMu.Unlock()
c.cfg.Code = code
c.kc.updateCode(code)
}
// ============================================================
// CGO wrapper methods (called with c.mu held)
// ============================================================
func (c *Client) cInit() error {
cPath := C.CString(c.cfg.DataPath)
cUserID := C.CString(c.cfg.UserID)
defer C.free(unsafe.Pointer(cPath))
defer C.free(unsafe.Pointer(cUserID))
ret := C.init(cPath, cUserID, nil)
return mapCError(int(ret))
}
func (c *Client) cEncryptData(corpID, staffID string, plaintext []byte) ([]byte, error) {
cCorpID := C.CString(corpID)
cStaffID := C.CString(staffID)
seqID := uuid.New().String()
cSeqID := C.CString(seqID)
defer C.free(unsafe.Pointer(cCorpID))
defer C.free(unsafe.Pointer(cStaffID))
defer C.free(unsafe.Pointer(cSeqID))
var outBuf *C.uchar
var outLen C.uint
ret := C.encryptData(
cCorpID,
cStaffID,
(*C.uchar)(unsafe.Pointer(&plaintext[0])),
C.uint(len(plaintext)),
nil, // id - NULL for group messages
&outBuf,
&outLen,
cSeqID,
(C.call_proxy_func)(C.goProxyBridge),
)
retCode := int(ret)
if retCode != cFunctionOK {
return nil, mapCError(retCode)
}
// Copy result to Go slice and free C buffer
result := C.GoBytes(unsafe.Pointer(outBuf), C.int(outLen))
C.freeCryptoBuf(unsafe.Pointer(outBuf))
return result, nil
}
func (c *Client) cDecryptData(corpID, staffID string, ciphertext []byte) ([]byte, error) {
cCorpID := C.CString(corpID)
cStaffID := C.CString(staffID)
seqID := uuid.New().String()
cSeqID := C.CString(seqID)
defer C.free(unsafe.Pointer(cCorpID))
defer C.free(unsafe.Pointer(cStaffID))
defer C.free(unsafe.Pointer(cSeqID))
var outBuf *C.uchar
var outLen C.uint
// The C layer treats the ciphertext as a NUL-terminated C string
// (sscanf/strtok). Go byte slices are not NUL-terminated and strtok would
// mutate the caller's buffer, so pass a NUL-terminated private copy.
cbuf := make([]byte, len(ciphertext)+1)
copy(cbuf, ciphertext)
ret := C.decryptData(
cCorpID,
cStaffID,
(*C.uchar)(unsafe.Pointer(&cbuf[0])),
C.uint(len(ciphertext)),
nil, // id
&outBuf,
&outLen,
cSeqID,
(C.call_proxy_func)(C.goProxyBridge),
)
retCode := int(ret)
if retCode != cFunctionOK {
return nil, mapCError(retCode)
}
result := C.GoBytes(unsafe.Pointer(outBuf), C.int(outLen))
C.freeCryptoBuf(unsafe.Pointer(outBuf))
return result, nil
}
func (c *Client) cEncryptFile(corpID, staffID, srcPath, dstPath string) error {
cCorpID := C.CString(corpID)
cStaffID := C.CString(staffID)
cID := C.CString("")
cSrcPath := C.CString(srcPath)
cDstPath := C.CString(dstPath)
seqID := uuid.New().String()
cSeqID := C.CString(seqID)
defer C.free(unsafe.Pointer(cCorpID))
defer C.free(unsafe.Pointer(cStaffID))
defer C.free(unsafe.Pointer(cID))
defer C.free(unsafe.Pointer(cSrcPath))
defer C.free(unsafe.Pointer(cDstPath))
defer C.free(unsafe.Pointer(cSeqID))
ret := C.encryptFile(
cCorpID,
cStaffID,
cID,
cSrcPath,
cDstPath,
cSeqID,
(C.call_proxy_func)(C.goProxyBridge),
)
retCode := int(ret)
if retCode != cFunctionOK {
return mapCError(retCode)
}
return nil
}
func (c *Client) cDecryptFile(corpID, staffID, srcPath, dstPath string) error {
cCorpID := C.CString(corpID)
cStaffID := C.CString(staffID)
cID := C.CString("")
cSrcPath := C.CString(srcPath)
cDstPath := C.CString(dstPath)
seqID := uuid.New().String()
cSeqID := C.CString(seqID)
defer C.free(unsafe.Pointer(cCorpID))
defer C.free(unsafe.Pointer(cStaffID))
defer C.free(unsafe.Pointer(cID))
defer C.free(unsafe.Pointer(cSrcPath))
defer C.free(unsafe.Pointer(cDstPath))
defer C.free(unsafe.Pointer(cSeqID))
ret := C.decryptFile(
cCorpID,
cStaffID,
cID,
cSrcPath,
cDstPath,
cSeqID,
(C.call_proxy_func)(C.goProxyBridge),
)
retCode := int(ret)
if retCode != cFunctionOK {
return mapCError(retCode)
}
return nil
}
func (c *Client) cEncryptBuffer(corpID, staffID string, data []byte) ([]byte, error) {
cCorpID := C.CString(corpID)
cStaffID := C.CString(staffID)
cID := C.CString("")
seqID := uuid.New().String()
cSeqID := C.CString(seqID)
defer C.free(unsafe.Pointer(cCorpID))
defer C.free(unsafe.Pointer(cStaffID))
defer C.free(unsafe.Pointer(cID))
defer C.free(unsafe.Pointer(cSeqID))
var outBuf *C.uchar
var outLen C.uint
ret := C.encryptBuffer(
cCorpID,
cStaffID,
(*C.uchar)(unsafe.Pointer(&data[0])),
C.uint(len(data)),
cID,
&outBuf,
&outLen,
cSeqID,
(C.call_proxy_func)(C.goProxyBridge),
)
retCode := int(ret)
if retCode != cFunctionOK {
return nil, mapCError(retCode)
}
result := C.GoBytes(unsafe.Pointer(outBuf), C.int(outLen))
C.freeCryptoBuf(unsafe.Pointer(outBuf))
return result, nil
}
func (c *Client) cDecryptBuffer(corpID, staffID string, data []byte) ([]byte, error) {
cCorpID := C.CString(corpID)
cStaffID := C.CString(staffID)
cID := C.CString("")
seqID := uuid.New().String()
cSeqID := C.CString(seqID)
defer C.free(unsafe.Pointer(cCorpID))
defer C.free(unsafe.Pointer(cStaffID))
defer C.free(unsafe.Pointer(cID))
defer C.free(unsafe.Pointer(cSeqID))
var outBuf *C.uchar
var outLen C.uint
ret := C.decryptBuffer(
cCorpID,
cStaffID,
(*C.uchar)(unsafe.Pointer(&data[0])),
C.uint(len(data)),
cID,
&outBuf,
&outLen,
cSeqID,
(C.call_proxy_func)(C.goProxyBridge),
)
retCode := int(ret)
if retCode != cFunctionOK {
return nil, mapCError(retCode)
}
result := C.GoBytes(unsafe.Pointer(outBuf), C.int(outLen))
C.freeCryptoBuf(unsafe.Pointer(outBuf))
return result, nil
}
func (c *Client) cSetResponse(corpID, jsonStr string) error {
cCorpID := C.CString(corpID)
cJSON := C.CString(jsonStr)
defer C.free(unsafe.Pointer(cCorpID))
defer C.free(unsafe.Pointer(cJSON))
ret := C.setResponse(
cCorpID,
cJSON,
(C.block_crypto_func)(C.goBlockBridge),
)
retCode := int(ret)
if retCode != cFunctionOK {
return mapCError(retCode)
}
return nil
}
func (c *Client) cSetPushData(corpID, staffID, pushData string) error {
cCorpID := C.CString(corpID)
cStaffID := C.CString(staffID)
cPushData := C.CString(pushData)
seqID := uuid.New().String()
cSeqID := C.CString(seqID)
defer C.free(unsafe.Pointer(cCorpID))
defer C.free(unsafe.Pointer(cStaffID))
defer C.free(unsafe.Pointer(cPushData))
defer C.free(unsafe.Pointer(cSeqID))
ret := C.setPushData(
cCorpID,
cStaffID,
cPushData,
cSeqID,
(C.call_proxy_func)(C.goProxyBridge),
(C.cancel_block_crypto_func)(C.goCancelBlockBridge),
)
retCode := int(ret)
if retCode != cFunctionOK && retCode != cSendRequestParamOK {
return mapCError(retCode)
}
return nil
}
+5
View File
@@ -0,0 +1,5 @@
package safechat
// Version is the current version of the SafeChat Go SDK.
// This version follows semantic versioning (semver).
const Version = "1.0.0"