Compare commits

..

3 Commits

Author SHA1 Message Date
evandance
40a0a9de66 feat(enhancement): centralize HTTP transport policies (#2021) 2026-08-02 14:55:05 +08:00
liangshuo-1
a8ad44ba13 docs: remove broken Star History chart (#2141) 2026-08-01 11:42:24 +08:00
liangshuo-1
003d0f42f8 chore: release v1.0.81 (#2136) 2026-07-31 18:47:19 +08:00
50 changed files with 4530 additions and 1666 deletions

View File

@@ -2,6 +2,35 @@
All notable changes to this project will be documented in this file.
## [v1.0.81] - 2026-07-31
### Features
- support visible_rule for form questions (#1891)
- **contact**: add bot search shortcut (#2083)
- add SXSD schema validation to Slides lint (#2103)
- **drive**: add comment-operation shortcuts (#1898)
- **drive**: extend permission shortcuts for Miaoda (#2070)
- **apps**: add cache debug commands (+cache-get/-delete/-clear) (#1896)
- support source file preview artifacts (#2085)
### Bug Fixes
- **contact**: stop bot match segments carrying tags or empty entries (#2115)
- **base**: resolve Base URL block types accurately (#2099)
- **drive**: use title for default download filename (#2089)
- drop stale target version from root upgrade prompt (#2100)
### Documentation
- **calendar**: warn against container-default timezone in time conversion (#2104)
- **calendar**: confirm scope before editing recurring events (#2119)
- **base**: clarify form and file operation routing (#2110)
### Misc
- add protected public domain allowlists (#2111)
## [v1.0.80] - 2026-07-29
### Features
@@ -1722,6 +1751,7 @@ Bundled AI agent skills for intelligent assistance:
- Bilingual documentation (English & Chinese).
- CI/CD pipelines: linting, testing, coverage reporting, and automated releases.
[v1.0.81]: https://github.com/larksuite/cli/releases/tag/v1.0.81
[v1.0.80]: https://github.com/larksuite/cli/releases/tag/v1.0.80
[v1.0.79]: https://github.com/larksuite/cli/releases/tag/v1.0.79
[v1.0.78]: https://github.com/larksuite/cli/releases/tag/v1.0.78

View File

@@ -310,10 +310,6 @@ lark-cli config risk-control default
Please fully understand all usage risks. By using this tool, you are deemed to voluntarily assume all related responsibilities.
## Star History
[![Star History Chart](https://api.star-history.com/svg?repos=larksuite/cli&type=Date)](https://star-history.com/#larksuite/cli&Date)
## Contributing
Community contributions are welcome! If you find a bug or have feature suggestions, please submit an [Issue](https://github.com/larksuite/cli/issues) or [Pull Request](https://github.com/larksuite/cli/pulls).

View File

@@ -311,10 +311,6 @@ lark-cli config risk-control default
请您充分知悉全部使用风险,使用本工具即视为您自愿承担相关所有责任。
## Star History
[![Star History Chart](https://api.star-history.com/svg?repos=larksuite/cli&type=Date)](https://star-history.com/#larksuite/cli&Date)
## 贡献
欢迎社区贡献!如果你发现 bug 或有功能建议,请提交 [Issue](https://github.com/larksuite/cli/issues) 或 [Pull Request](https://github.com/larksuite/cli/pulls)。

View File

@@ -179,8 +179,8 @@ func runCreateAppFlow(ctx context.Context, f *cmdutil.Factory, brandOverride cor
}
// Step 1: Request app registration (begin)
// Use the shared proxy-plugin-aware transport so registration traffic is not
// a bypass of proxy plugin mode.
// Registration is platform traffic, so it must use the provider-aware
// transport as well as the shared proxy configuration.
httpClient := transport.NewHTTPClient(0)
authResp, err := larkauth.RequestAppRegistration(ctx, httpClient, larkBrand, f.IOStreams.ErrOut)
if err != nil {

View File

@@ -157,8 +157,8 @@ func networkChecks(ctx context.Context, opts *DoctorOptions, ep core.Endpoints)
}
}
// Use the shared proxy-plugin-aware transport so connectivity checks reflect
// the real egress path (and are blocked when proxy plugin fails closed).
// Connectivity checks are platform traffic and must exercise the same
// provider-aware route as real platform requests.
httpClient := transport.NewHTTPClient(0)
mcpURL := ep.MCP + "/mcp"

View File

@@ -12,9 +12,18 @@ import (
"net/http"
"testing"
exttransport "github.com/larksuite/cli/extension/transport"
"github.com/larksuite/cli/internal/envvars"
internaltransport "github.com/larksuite/cli/internal/transport"
"github.com/larksuite/cli/sidecar"
)
type roundTripFunc func(*http.Request) (*http.Response, error)
func (f roundTripFunc) RoundTrip(req *http.Request) (*http.Response, error) {
return f(req)
}
// failingBody is a ReadCloser that errors on Read and tracks Close calls.
type failingBody struct {
err error
@@ -263,3 +272,55 @@ func TestInterceptor_EmptyBody(t *testing.T) {
t.Errorf("body SHA256 = %q, want empty-string SHA256 %q", sha, expectedEmpty)
}
}
func TestLegacySidecarProviderStillHandlesForcedExternalRequests(t *testing.T) {
t.Setenv(envvars.CliAuthProxy, "http://127.0.0.1:16384")
t.Setenv(envvars.CliProxyKey, "test-key")
previousProvider := exttransport.GetProvider()
exttransport.Register(&Provider{})
t.Cleanup(func() { exttransport.Register(previousProvider) })
seen := make(chan *http.Request, 2)
base := roundTripFunc(func(req *http.Request) (*http.Response, error) {
seen <- req.Clone(req.Context())
return &http.Response{StatusCode: http.StatusNoContent, Body: http.NoBody, Request: req}, nil
})
client := internaltransport.ClientForRequestClass(
&http.Client{Transport: internaltransport.NewHTTPPolicyRouter(base, base)},
exttransport.RequestClassExternal,
)
withSentinel, err := http.NewRequest(http.MethodGet, "https://external.example/protected", nil)
if err != nil {
t.Fatal(err)
}
withSentinel.Header.Set("Authorization", "Bearer "+sidecar.SentinelUAT)
resp, err := client.Do(withSentinel)
if err != nil {
t.Fatal(err)
}
resp.Body.Close()
withoutSentinel, err := http.NewRequest(http.MethodGet, "https://external.example/public", nil)
if err != nil {
t.Fatal(err)
}
resp, err = client.Do(withoutSentinel)
if err != nil {
t.Fatal(err)
}
resp.Body.Close()
proxied := <-seen
if proxied.URL.Scheme != "http" || proxied.URL.Host != "127.0.0.1:16384" {
t.Fatalf("sentinel request URL = %s, want sidecar route", proxied.URL)
}
if got := proxied.Header.Get(sidecar.HeaderProxyTarget); got != "https://external.example" {
t.Fatalf("sentinel request proxy target = %q", got)
}
passthrough := <-seen
if got := passthrough.URL.String(); got != "https://external.example/public" {
t.Fatalf("non-sentinel request URL = %q, want unchanged", got)
}
}

View File

@@ -15,6 +15,27 @@ type Provider interface {
ResolveInterceptor(ctx context.Context) Interceptor
}
// RequestClass describes the trust boundary of an outbound HTTP request.
// Platform requests target endpoints owned by the CLI's endpoint resolver;
// external requests target user-provided, pre-signed, CDN, registry, or other
// non-platform URLs. Redirect targets are classified again from each hop's
// logical URL; rewriting a host in an interceptor does not add that host to
// the platform endpoint catalog.
type RequestClass string
const (
RequestClassPlatform RequestClass = "platform"
RequestClassExternal RequestClass = "external"
)
// ScopedProvider optionally limits a Provider to selected request classes.
// Providers that do not implement this interface retain the original
// behavior and apply to every request class.
type ScopedProvider interface {
Provider
SupportsRequestClass(RequestClass) bool
}
// Interceptor defines network-layer customization via a pre/post hook pair.
// The built-in transport chain always executes between PreRoundTrip and the
// returned post function, and cannot be skipped or overridden by the extension.

View File

@@ -17,6 +17,8 @@ import (
"github.com/larksuite/cli/internal/transport"
)
var _ transport.RoundTripperDecorator = (*SecurityPolicyTransport)(nil)
// SecurityPolicyTransport is an http.RoundTripper that intercepts all responses
// and checks for security policy errors.
type SecurityPolicyTransport struct {
@@ -31,6 +33,16 @@ func (t *SecurityPolicyTransport) base() http.RoundTripper {
return transport.Fallback()
}
func (t *SecurityPolicyTransport) BaseRoundTripper() http.RoundTripper {
return t.base()
}
func (t *SecurityPolicyTransport) WithBaseRoundTripper(base http.RoundTripper) http.RoundTripper {
cloned := *t
cloned.Base = base
return &cloned
}
// RoundTrip implements http.RoundTripper.
func (t *SecurityPolicyTransport) RoundTrip(req *http.Request) (*http.Response, error) {
resp, err := t.base().RoundTrip(req)

View File

@@ -212,6 +212,9 @@ func (c *APIClient) DoStream(ctx context.Context, req *larkcore.ApiReq, as core.
resp, err := httpClient.Do(httpReq)
if err != nil {
cancel()
if _, ok := errs.ProblemOf(err); ok {
return nil, err
}
return nil, errs.NewNetworkError(classifyNetworkSubtype(err), "stream request failed: %s", err).WithCause(err)
}
resp.Body = &cancelOnCloseBody{ReadCloser: resp.Body, cancel: cancel}

View File

@@ -518,6 +518,29 @@ func TestDoStream_TransportFailureSplitsSubtype(t *testing.T) {
}
}
func TestDoStream_PreservesTypedTransportError(t *testing.T) {
policyErr := errs.NewSecurityPolicyError(errs.SubtypeAccessDenied, "blocked redirect")
ac := &APIClient{
HTTP: &http.Client{Transport: roundTripFunc(func(*http.Request) (*http.Response, error) {
return nil, policyErr
})},
Credential: credential.NewCredentialProvider(nil, nil, &staticTokenResolver{}, nil),
Config: &core.CliConfig{AppID: "test-app", AppSecret: "test-secret", Brand: core.BrandFeishu},
}
_, err := ac.DoStream(context.Background(), &larkcore.ApiReq{
HttpMethod: http.MethodGet,
ApiPath: "/open-apis/drive/v1/files/file_token/download",
}, core.AsBot)
problem, ok := errs.ProblemOf(err)
if !ok || problem.Category != errs.CategoryPolicy || problem.Subtype != errs.SubtypeAccessDenied {
t.Fatalf("DoStream() problem = %#v, %v; want policy/access_denied", problem, ok)
}
if !errors.Is(err, policyErr) {
t.Fatal("DoStream() did not preserve the typed transport error")
}
}
// failingTokenResolver always returns TokenUnavailableError, exercising the
// auth/credential failure path through resolveAccessToken.
type failingTokenResolver struct{}

View File

@@ -16,10 +16,12 @@ import (
"github.com/larksuite/cli/errs"
extcred "github.com/larksuite/cli/extension/credential"
"github.com/larksuite/cli/extension/fileio"
exttransport "github.com/larksuite/cli/extension/transport"
"github.com/larksuite/cli/internal/client"
"github.com/larksuite/cli/internal/core"
"github.com/larksuite/cli/internal/credential"
"github.com/larksuite/cli/internal/keychain"
"github.com/larksuite/cli/internal/transport"
)
// Factory holds shared dependencies injected into every command.
@@ -31,7 +33,7 @@ type InvocationContext struct {
type Factory struct {
Config func() (*core.CliConfig, error) // lazily loads app config from Credential
HttpClient func() (*http.Client, error) // HTTP client for non-Lark API calls (with retry and security headers)
HttpClient func() (*http.Client, error) // policy-routed HTTP client for direct requests
LarkClient func() (*lark.Client, error) // Lark SDK client for all Open API calls
IOStreams *IOStreams // stdin/stdout/stderr streams
@@ -48,6 +50,18 @@ type Factory struct {
SkillContent fs.FS // embedded skill tree (rooted at the skill list); nil when the build embeds no skills
}
// ExternalHTTPClient returns a clone of the existing Factory client whose
// requests are explicitly classified as external. The underlying client,
// redirect policy, timeout, proxy configuration, and legacy transport provider
// behavior are preserved.
func (f *Factory) ExternalHTTPClient() (*http.Client, error) {
client, err := f.HttpClient()
if err != nil {
return nil, err
}
return transport.ClientForRequestClass(client, exttransport.RequestClassExternal), nil
}
// ResolveFileIO resolves a FileIO instance using the current execution context.
// The provider controls whether the returned instance is fresh or cached.
func (f *Factory) ResolveFileIO(ctx context.Context) fileio.FileIO {

View File

@@ -5,16 +5,18 @@ package cmdutil
import (
"context"
"fmt"
"io"
"net/http"
"net/url"
"os"
"strings"
"sync"
"time"
lark "github.com/larksuite/oapi-sdk-go/v3"
larkcore "github.com/larksuite/oapi-sdk-go/v3/core"
"github.com/larksuite/cli/errs"
extcred "github.com/larksuite/cli/extension/credential"
"github.com/larksuite/cli/extension/fileio"
"github.com/larksuite/cli/internal/auth"
@@ -48,6 +50,19 @@ func NewDefault(streams *IOStreams, inv InvocationContext) *Factory {
// workspace-scoped. Default is WorkspaceLocal — existing behavior unchanged.
ws := core.DetectWorkspaceFromEnv(os.Getenv)
core.SetCurrentWorkspace(ws)
workspaceConfig := core.NewConfigSnapshot()
bootstrapHostSignalSource := sync.OnceValue(func() riskcontrol.Source {
return resolveSDKHostSignalSource(workspaceConfig)
})
// Install after workspace selection so the dependency bootstrap bridge uses
// the correct shared proxy configuration. NewDefault is also used by cmd.Build
// consumers, so this keeps their request routing identical to cmd.Execute.
transport.InstallSDKTransportBridge(func(base http.RoundTripper) http.RoundTripper {
return buildSDKPlatformTransportWithBase(
base,
bootstrapHostSignalSource(),
)
})
// Inject workspace-aware dir into keychain's log system.
// This breaks the core↔keychain import cycle by using a function variable.
@@ -55,7 +70,6 @@ func NewDefault(streams *IOStreams, inv InvocationContext) *Factory {
// Phase 0: FileIO provider (no dependency)
f.FileIOProvider = fileio.GetProvider()
workspaceConfig := core.NewConfigSnapshot()
// Phase 1: HttpClient (no credential dependency)
f.HttpClient = cachedHttpClientFunc(f, workspaceConfig)
@@ -87,15 +101,45 @@ func NewDefault(streams *IOStreams, inv InvocationContext) *Factory {
return f
}
// safeRedirectPolicy prevents credential headers from being forwarded
// when a response redirects to a different host (e.g. Lark API 302 → CDN).
// Strips Authorization, X-Lark-MCP-UAT, and X-Lark-MCP-TAT on cross-host
// redirects; other headers like X-Cli-* pass through.
// safeRedirectPolicy permits cross-origin redirects only for bodyless GET and
// HEAD requests. This allows API download redirects while preventing OAuth or
// other credential-bearing request bodies from being replayed to another
// origin. HTTPS requests can never be downgraded to HTTP.
func safeRedirectPolicy(req *http.Request, via []*http.Request) error {
if len(via) >= 10 {
return fmt.Errorf("too many redirects")
return errs.NewNetworkError(errs.SubtypeNetworkTransport, "too many redirects")
}
if len(via) > 0 && req.URL.Host != via[0].URL.Host {
if len(via) == 0 {
return nil
}
original := via[0]
previous := via[len(via)-1]
if previous.URL != nil && req.URL != nil && strings.EqualFold(previous.URL.Scheme, "https") && !strings.EqualFold(req.URL.Scheme, "https") {
return errs.NewSecurityPolicyError(
errs.SubtypeAccessDenied,
"redirect from HTTPS to %s is not allowed",
req.URL.Scheme,
)
}
if !sameRedirectOrigin(previous.URL, req.URL) {
if req.Method != http.MethodGet && req.Method != http.MethodHead {
return errs.NewSecurityPolicyError(
errs.SubtypeAccessDenied,
"cross-origin redirect for HTTP method %s is not allowed",
req.Method,
)
}
if req.Body != nil || req.GetBody != nil {
return errs.NewSecurityPolicyError(
errs.SubtypeAccessDenied,
"cross-origin redirect with a request body is not allowed",
)
}
}
// net/http copies initial headers onto every redirect request. Continue
// stripping credentials for every hop outside the initial origin, even when
// two consecutive redirect targets share an origin.
if !sameRedirectOrigin(original.URL, req.URL) {
req.Header.Del("Authorization")
req.Header.Del("X-Lark-MCP-UAT")
req.Header.Del("X-Lark-MCP-TAT")
@@ -103,6 +147,29 @@ func safeRedirectPolicy(req *http.Request, via []*http.Request) error {
return nil
}
func sameRedirectOrigin(left, right *url.URL) bool {
if left == nil || right == nil {
return false
}
return strings.EqualFold(left.Scheme, right.Scheme) &&
strings.EqualFold(left.Hostname(), right.Hostname()) &&
effectivePort(left) == effectivePort(right)
}
func effectivePort(candidate *url.URL) string {
if port := candidate.Port(); port != "" {
return port
}
switch strings.ToLower(candidate.Scheme) {
case "http":
return "80"
case "https":
return "443"
default:
return ""
}
}
// warnIfProxied is a test seam for the proxy-warning gate. Production wires it
// to transport.WarnIfProxied; tests swap in a spy to count invocations. It is
// needed because the real function is guarded by an internal sync.Once, so
@@ -118,15 +185,12 @@ func cachedHttpClientFunc(f *Factory, workspaceConfig workspaceConfigSource) fun
}
hostSignalSource := resolveSDKHostSignalSource(workspaceConfig)
var rt http.RoundTripper = transport.Shared()
rt = riskcontrol.NewTransport(rt, hostSignalSource)
rt = &RetryTransport{Base: rt}
rt = &SecurityHeaderTransport{Base: rt}
rt = &auth.SecurityPolicyTransport{Base: rt} // Add our global response interceptor
rt = wrapWithExtension(rt)
shared := transport.Shared()
outbound := riskcontrol.NewTransport(shared, hostSignalSource)
platform := buildDirectHTTPTransport(outbound, true)
external := buildDirectHTTPTransport(outbound, false)
client := &http.Client{
Transport: rt,
Transport: transport.NewHTTPPolicyRouter(platform, external),
Timeout: 30 * time.Second,
CheckRedirect: safeRedirectPolicy,
}
@@ -134,6 +198,15 @@ func cachedHttpClientFunc(f *Factory, workspaceConfig workspaceConfigSource) fun
})
}
func buildDirectHTTPTransport(base http.RoundTripper, platform bool) http.RoundTripper {
var builtIn http.RoundTripper = &RetryTransport{Base: base}
builtIn = &SecurityHeaderTransport{Base: builtIn}
if platform {
builtIn = &auth.SecurityPolicyTransport{Base: builtIn}
}
return builtIn
}
func cachedLarkClientFunc(f *Factory, workspaceConfig workspaceConfigSource) func() (*lark.Client, error) {
return sync.OnceValues(func() (*lark.Client, error) {
acct, err := f.Credential.ResolveAccount(context.Background())
@@ -149,14 +222,8 @@ func cachedLarkClientFunc(f *Factory, workspaceConfig workspaceConfigSource) fun
warnIfProxied(f.IOStreams.ErrOut)
}
hostSignalSource := resolveSDKHostSignalSource(workspaceConfig)
var sdkBase http.RoundTripper = transport.Shared()
// The innermost SDK boundary always strips reserved host-signal headers;
// a nil source makes it strip-only when workspace policy disables signal
// collection.
sdkBase = riskcontrol.NewTransport(sdkBase, hostSignalSource)
sdkTransport := wrapSDKTransport(sdkBase)
opts = append(opts, lark.WithHttpClient(&http.Client{
Transport: sdkTransport,
Transport: buildSDKTransport(hostSignalSource),
CheckRedirect: safeRedirectPolicy,
}))
ep := core.ResolveEndpoints(acct.Brand)
@@ -165,12 +232,41 @@ func cachedLarkClientFunc(f *Factory, workspaceConfig workspaceConfigSource) fun
})
}
func wrapSDKTransport(next http.RoundTripper) http.RoundTripper {
var sdkTransport http.RoundTripper = &RetryTransport{Base: next}
sdkTransport = &UserAgentTransport{Base: sdkTransport}
sdkTransport = &BuildHeaderTransport{Base: sdkTransport}
sdkTransport = &auth.SecurityPolicyTransport{Base: sdkTransport}
return wrapWithExtension(sdkTransport)
func buildSDKTransport(hostSignalSource riskcontrol.Source) http.RoundTripper {
return buildSDKTransportWithBase(transport.Shared(), hostSignalSource)
}
func buildSDKPlatformTransportWithBase(
base http.RoundTripper,
hostSignalSource riskcontrol.Source,
) http.RoundTripper {
outbound := riskcontrol.NewTransport(base, hostSignalSource)
return buildSDKHTTPTransport(outbound, true)
}
func buildSDKTransportWithBase(
base http.RoundTripper,
hostSignalSource riskcontrol.Source,
) http.RoundTripper {
// Risk control is the innermost trusted boundary for both request classes.
// It therefore observes the final URL and strips extension-supplied reserved
// headers immediately before the network transport.
outbound := riskcontrol.NewTransport(base, hostSignalSource)
return transport.NewHTTPPolicyRouter(
buildSDKHTTPTransport(outbound, true),
buildSDKHTTPTransport(outbound, false),
)
}
func buildSDKHTTPTransport(base http.RoundTripper, platform bool) http.RoundTripper {
var builtIn http.RoundTripper = &RetryTransport{Base: base}
builtIn = &UserAgentTransport{Base: builtIn}
builtIn = &BuildHeaderTransport{Base: builtIn}
builtIn = &SecurityHeaderTransport{Base: builtIn}
if platform {
builtIn = &auth.SecurityPolicyTransport{Base: builtIn}
}
return builtIn
}
type credentialDeps struct {

View File

@@ -4,13 +4,20 @@
package cmdutil
import (
"errors"
"io"
"net/http"
"net/http/httptest"
"strings"
"testing"
"github.com/larksuite/cli/errs"
exttransport "github.com/larksuite/cli/extension/transport"
"github.com/larksuite/cli/internal/core"
internaltransport "github.com/larksuite/cli/internal/transport"
)
func TestCachedHttpClientFunc_ReturnsSameInstance(t *testing.T) {
func TestCachedHTTPClientFunc_ReturnsSameInstance(t *testing.T) {
isEnabled := false
f, _, _, _ := TestFactory(t, &core.CliConfig{AppID: "test-app"})
f.IOStreams.ErrOut = io.Discard
@@ -33,7 +40,7 @@ func TestCachedHttpClientFunc_ReturnsSameInstance(t *testing.T) {
}
}
func TestCachedHttpClientFunc_HasTimeout(t *testing.T) {
func TestCachedHTTPClientFunc_HasTimeout(t *testing.T) {
isEnabled := false
f, _, _, _ := TestFactory(t, &core.CliConfig{AppID: "test-app"})
f.IOStreams.ErrOut = io.Discard
@@ -44,7 +51,7 @@ func TestCachedHttpClientFunc_HasTimeout(t *testing.T) {
}
}
func TestCachedHttpClientFunc_HasRedirectPolicy(t *testing.T) {
func TestCachedHTTPClientFunc_HasRedirectPolicy(t *testing.T) {
isEnabled := false
f, _, _, _ := TestFactory(t, &core.CliConfig{AppID: "test-app"})
f.IOStreams.ErrOut = io.Discard
@@ -54,3 +61,283 @@ func TestCachedHttpClientFunc_HasRedirectPolicy(t *testing.T) {
t.Error("expected CheckRedirect to be set (safeRedirectPolicy)")
}
}
func TestFactoryExternalHTTPClientClonesExistingClient(t *testing.T) {
base := &http.Client{Timeout: 17, CheckRedirect: safeRedirectPolicy}
factory := &Factory{HttpClient: func() (*http.Client, error) { return base, nil }}
external, err := factory.ExternalHTTPClient()
if err != nil {
t.Fatal(err)
}
if external == base {
t.Fatal("ExternalHTTPClient returned the cached client instead of a clone")
}
if external.Timeout != base.Timeout || external.CheckRedirect == nil {
t.Fatal("ExternalHTTPClient did not preserve client policy")
}
if base.Transport != nil {
t.Fatal("ExternalHTTPClient mutated the cached client's transport")
}
}
type platformOnlyStubProvider struct {
*stubTransportProvider
}
func (*platformOnlyStubProvider) SupportsRequestClass(class exttransport.RequestClass) bool {
return class == exttransport.RequestClassPlatform
}
func TestFactoryHTTPClientRoutesPoliciesByRequestClass(t *testing.T) {
t.Setenv("LARKSUITE_CLI_NO_PROXY", "1")
interceptor := &headerCapturingInterceptor{}
exttransport.Register(&platformOnlyStubProvider{stubTransportProvider: &stubTransportProvider{interceptor: interceptor}})
t.Cleanup(func() { exttransport.Register(nil) })
received := make(chan http.Header, 2)
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, req *http.Request) {
received <- req.Header.Clone()
w.WriteHeader(http.StatusNoContent)
}))
t.Cleanup(server.Close)
factory := &Factory{IOStreams: &IOStreams{ErrOut: io.Discard}}
client, err := cachedHttpClientFunc(factory, nil)()
if err != nil {
t.Fatal(err)
}
factory.HttpClient = func() (*http.Client, error) { return client, nil }
platformClient := internaltransport.ClientForRequestClass(client, exttransport.RequestClassPlatform)
externalClient, err := factory.ExternalHTTPClient()
if err != nil {
t.Fatal(err)
}
for _, client := range []*http.Client{platformClient, externalClient} {
resp, err := client.Get(server.URL)
if err != nil {
t.Fatal(err)
}
resp.Body.Close()
}
platformHeaders := <-received
if got := platformHeaders.Get("X-Custom-Trace"); got != "ext-trace-123" {
t.Fatalf("platform extension header = %q, want ext-trace-123", got)
}
if got := platformHeaders.Get(HeaderSource); got != SourceValue {
t.Fatalf("platform security header = %q, want %q", got, SourceValue)
}
externalHeaders := <-received
if got := externalHeaders.Get("X-Custom-Trace"); got != "" {
t.Fatalf("external request leaked extension header %q", got)
}
for header, values := range BaseSecurityHeaders() {
if len(values) == 0 {
continue
}
want := values[len(values)-1]
if got := externalHeaders.Get(header); got != want {
t.Fatalf("external security header %s = %q, want preserved value %q", header, got, want)
}
}
}
func TestFactoryExternalHTTPClientDoesNotParsePlatformErrorProtocol(t *testing.T) {
t.Setenv("LARKSUITE_CLI_NO_PROXY", "1")
exttransport.Register(nil)
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, req *http.Request) {
w.Header().Set("Content-Type", "application/json")
_, _ = io.WriteString(w, `{"code":21000,"msg":"application-defined external response","data":{"cli_hint":"external-defined"}}`)
}))
t.Cleanup(server.Close)
factory := &Factory{IOStreams: &IOStreams{ErrOut: io.Discard}}
client, err := cachedHttpClientFunc(factory, nil)()
if err != nil {
t.Fatal(err)
}
factory.HttpClient = func() (*http.Client, error) { return client, nil }
platform := internaltransport.ClientForRequestClass(client, exttransport.RequestClassPlatform)
if _, err := platform.Get(server.URL); err == nil {
t.Fatal("platform request error = nil, want security policy classification")
} else {
var policyErr *errs.SecurityPolicyError
if !errors.As(err, &policyErr) {
t.Fatalf("platform request error type = %T, want *errs.SecurityPolicyError", err)
}
}
external, err := factory.ExternalHTTPClient()
if err != nil {
t.Fatal(err)
}
resp, err := external.Get(server.URL)
if err != nil {
t.Fatalf("external request parsed platform error protocol: %v", err)
}
resp.Body.Close()
}
func TestSafeRedirectPolicyAllowsBodylessCrossOriginGetAndStripsCredentials(t *testing.T) {
original, err := http.NewRequest(http.MethodGet, "https://open.feishu.cn/start", nil)
if err != nil {
t.Fatal(err)
}
redirect, err := http.NewRequest(http.MethodGet, "https://cdn.example.com/file", nil)
if err != nil {
t.Fatal(err)
}
for _, header := range []string{"Authorization", "X-Lark-MCP-UAT", "X-Lark-MCP-TAT"} {
redirect.Header.Set(header, "secret")
}
if err := safeRedirectPolicy(redirect, []*http.Request{original}); err != nil {
t.Fatalf("safeRedirectPolicy() error = %v, want allowed GET redirect", err)
}
for _, header := range []string{"Authorization", "X-Lark-MCP-UAT", "X-Lark-MCP-TAT"} {
if got := redirect.Header.Get(header); got != "" {
t.Fatalf("redirect retained %s=%q", header, got)
}
}
}
func TestSafeRedirectPolicyRejectsHTTPSDowngrade(t *testing.T) {
original, err := http.NewRequest(http.MethodGet, "https://open.feishu.cn/start", nil)
if err != nil {
t.Fatal(err)
}
redirect, err := http.NewRequest(http.MethodGet, "http://open.feishu.cn/next", nil)
if err != nil {
t.Fatal(err)
}
err = safeRedirectPolicy(redirect, []*http.Request{original})
if err == nil || !strings.Contains(err.Error(), "HTTPS") {
t.Fatalf("safeRedirectPolicy() error = %v, want HTTPS downgrade rejection", err)
}
requireRedirectProblem(t, err, errs.CategoryPolicy, errs.SubtypeAccessDenied)
}
func TestSafeRedirectPolicyRejectsCrossOriginMethod(t *testing.T) {
original, err := http.NewRequest(http.MethodPost, "https://accounts.feishu.cn/token", nil)
if err != nil {
t.Fatal(err)
}
redirect, err := http.NewRequest(http.MethodPost, "https://external.example/token", nil)
if err != nil {
t.Fatal(err)
}
err = safeRedirectPolicy(redirect, []*http.Request{original})
if err == nil || !strings.Contains(err.Error(), "HTTP method POST") {
t.Fatalf("safeRedirectPolicy() error = %v, want cross-origin method rejection", err)
}
requireRedirectProblem(t, err, errs.CategoryPolicy, errs.SubtypeAccessDenied)
}
func TestSafeRedirectPolicyRejectsCrossOriginRequestBody(t *testing.T) {
original, err := http.NewRequest(http.MethodGet, "https://accounts.feishu.cn/token", nil)
if err != nil {
t.Fatal(err)
}
redirect, err := http.NewRequest(http.MethodGet, "https://external.example/token", strings.NewReader("client_secret=secret"))
if err != nil {
t.Fatal(err)
}
err = safeRedirectPolicy(redirect, []*http.Request{original})
if err == nil || !strings.Contains(err.Error(), "request body") {
t.Fatalf("safeRedirectPolicy() error = %v, want cross-origin body rejection", err)
}
requireRedirectProblem(t, err, errs.CategoryPolicy, errs.SubtypeAccessDenied)
}
func TestSafeRedirectPolicyRejectsTooManyRedirects(t *testing.T) {
err := safeRedirectPolicy(&http.Request{}, make([]*http.Request, 10))
if err == nil || err.Error() != "too many redirects" {
t.Fatalf("safeRedirectPolicy() error = %v, want redirect limit rejection", err)
}
requireRedirectProblem(t, err, errs.CategoryNetwork, errs.SubtypeNetworkTransport)
}
func TestSafeRedirectPolicyTreatsDefaultHTTPSPortAsSameOrigin(t *testing.T) {
original, err := http.NewRequest(http.MethodPost, "https://accounts.feishu.cn/token", strings.NewReader("secret"))
if err != nil {
t.Fatal(err)
}
redirect, err := http.NewRequest(http.MethodPost, "https://accounts.feishu.cn:443/token-next", strings.NewReader("secret"))
if err != nil {
t.Fatal(err)
}
if err := safeRedirectPolicy(redirect, []*http.Request{original}); err != nil {
t.Fatalf("safeRedirectPolicy() error = %v, want same-origin redirect", err)
}
}
func TestSafeRedirectPolicyKeepsCredentialsStrippedAcrossExternalHops(t *testing.T) {
original, err := http.NewRequest(http.MethodGet, "https://open.feishu.cn/start", nil)
if err != nil {
t.Fatal(err)
}
previous, err := http.NewRequest(http.MethodGet, "https://cdn.example.com/first", nil)
if err != nil {
t.Fatal(err)
}
redirect, err := http.NewRequest(http.MethodGet, "https://cdn.example.com/second", nil)
if err != nil {
t.Fatal(err)
}
redirect.Header.Set("Authorization", "Bearer copied-from-initial-request")
if err := safeRedirectPolicy(redirect, []*http.Request{original, previous}); err != nil {
t.Fatalf("safeRedirectPolicy() error = %v, want same-CDN redirect", err)
}
if got := redirect.Header.Get("Authorization"); got != "" {
t.Fatalf("redirect retained Authorization=%q outside the initial origin", got)
}
}
func TestSafeRedirectPolicyRejectsDowngradeOnLaterHop(t *testing.T) {
original, err := http.NewRequest(http.MethodGet, "http://source.example/start", nil)
if err != nil {
t.Fatal(err)
}
previous, err := http.NewRequest(http.MethodGet, "https://cdn.example.com/secure", nil)
if err != nil {
t.Fatal(err)
}
redirect, err := http.NewRequest(http.MethodGet, "http://cdn.example.com/plain", nil)
if err != nil {
t.Fatal(err)
}
err = safeRedirectPolicy(redirect, []*http.Request{original, previous})
if err == nil || !strings.Contains(err.Error(), "HTTPS") {
t.Fatalf("safeRedirectPolicy() error = %v, want later-hop HTTPS downgrade rejection", err)
}
requireRedirectProblem(t, err, errs.CategoryPolicy, errs.SubtypeAccessDenied)
}
func requireRedirectProblem(t *testing.T, err error, category errs.Category, subtype errs.Subtype) {
t.Helper()
problem, ok := errs.ProblemOf(err)
if !ok {
t.Fatalf("error type = %T, want typed error", err)
}
if problem.Category != category || problem.Subtype != subtype {
t.Fatalf(
"error category/subtype = %s/%s, want %s/%s",
problem.Category,
problem.Subtype,
category,
subtype,
)
}
}

View File

@@ -34,9 +34,9 @@ var proxyWarnGateCases = []struct {
{"non-terminal stderr stays silent", false, 0},
}
// TestCachedHttpClientFunc_ProxyWarnGate verifies the http-client init path
// TestCachedHTTPClientFunc_ProxyWarnGate verifies the HTTP client init path
// invokes WarnIfProxied only when stderr is an interactive terminal.
func TestCachedHttpClientFunc_ProxyWarnGate(t *testing.T) {
func TestCachedHTTPClientFunc_ProxyWarnGate(t *testing.T) {
isEnabled := false
for _, tc := range proxyWarnGateCases {
t.Run(tc.name, func(t *testing.T) {

View File

@@ -46,7 +46,7 @@ func TestTestFactory_ReplacesGlobals(t *testing.T) {
URL: "/test",
Body: "ok",
})
// Use the stub via Factory HttpClient
// Use the stub via Factory HttpClient.
httpClient, err := f.HttpClient()
if err != nil {
t.Fatalf("HttpClient() error: %v", err)

View File

@@ -4,14 +4,19 @@
package cmdutil
import (
"context"
"net/http"
"time"
exttransport "github.com/larksuite/cli/extension/transport"
"github.com/larksuite/cli/internal/transport"
)
var (
_ transport.RoundTripperDecorator = (*RetryTransport)(nil)
_ transport.RoundTripperDecorator = (*UserAgentTransport)(nil)
_ transport.RoundTripperDecorator = (*BuildHeaderTransport)(nil)
_ transport.RoundTripperDecorator = (*SecurityHeaderTransport)(nil)
)
// RetryTransport is an http.RoundTripper that retries on 5xx responses
// and network errors. MaxRetries defaults to 0 (no retries).
type RetryTransport struct {
@@ -27,6 +32,16 @@ func (t *RetryTransport) base() http.RoundTripper {
return transport.Fallback()
}
func (t *RetryTransport) BaseRoundTripper() http.RoundTripper {
return t.base()
}
func (t *RetryTransport) WithBaseRoundTripper(base http.RoundTripper) http.RoundTripper {
cloned := *t
cloned.Base = base
return &cloned
}
func (t *RetryTransport) delay() time.Duration {
if t.Delay > 0 {
return t.Delay
@@ -63,6 +78,19 @@ type UserAgentTransport struct {
Base http.RoundTripper
}
func (t *UserAgentTransport) BaseRoundTripper() http.RoundTripper {
if t.Base != nil {
return t.Base
}
return transport.Fallback()
}
func (t *UserAgentTransport) WithBaseRoundTripper(base http.RoundTripper) http.RoundTripper {
cloned := *t
cloned.Base = base
return &cloned
}
func (t *UserAgentTransport) RoundTrip(req *http.Request) (*http.Response, error) {
req = req.Clone(req.Context())
req.Header.Set(HeaderUserAgent, UserAgentValue())
@@ -73,14 +101,25 @@ func (t *UserAgentTransport) RoundTrip(req *http.Request) (*http.Response, error
}
// BuildHeaderTransport is an http.RoundTripper that force-writes the
// X-Cli-Build header before every request. Used in the SDK transport chain,
// where SecurityHeaderTransport is not installed, to prevent extensions from
// tampering with the build classification. The direct HTTP chain is already
// covered by SecurityHeaderTransport iterating BaseSecurityHeaders.
// X-Cli-Build header before every request. It remains in the SDK transport
// chain as a narrow defense-in-depth layer alongside SecurityHeaderTransport.
type BuildHeaderTransport struct {
Base http.RoundTripper
}
func (t *BuildHeaderTransport) BaseRoundTripper() http.RoundTripper {
if t.Base != nil {
return t.Base
}
return transport.Fallback()
}
func (t *BuildHeaderTransport) WithBaseRoundTripper(base http.RoundTripper) http.RoundTripper {
cloned := *t
cloned.Base = base
return &cloned
}
func (t *BuildHeaderTransport) RoundTrip(req *http.Request) (*http.Response, error) {
req = req.Clone(req.Context())
req.Header.Set(HeaderBuild, DetectBuildKind())
@@ -103,6 +142,16 @@ func (t *SecurityHeaderTransport) base() http.RoundTripper {
return transport.Fallback()
}
func (t *SecurityHeaderTransport) BaseRoundTripper() http.RoundTripper {
return t.base()
}
func (t *SecurityHeaderTransport) WithBaseRoundTripper(base http.RoundTripper) http.RoundTripper {
cloned := *t
cloned.Base = base
return &cloned
}
// RoundTrip implements http.RoundTripper.
func (t *SecurityHeaderTransport) RoundTrip(req *http.Request) (*http.Response, error) {
req = req.Clone(req.Context())
@@ -120,67 +169,3 @@ func (t *SecurityHeaderTransport) RoundTrip(req *http.Request) (*http.Response,
}
return t.base().RoundTrip(req)
}
// extensionMiddleware wraps the built-in transport chain with pre/post hooks.
// The built-in chain always executes unless the extension is an
// exttransport.AbortableInterceptor and its PreRoundTripE returns a non-nil
// error; it cannot otherwise be skipped or overridden.
//
// The original request context is restored after the pre hook to prevent
// extensions from tampering with cancellation, deadlines, or built-in values.
// Cloning the request isolates header/URL/etc. mutations from the caller's
// request object; req.Body is intentionally shared — extensions that consume
// it are responsible for rewinding (see Interceptor doc).
type extensionMiddleware struct {
Base http.RoundTripper
Ext exttransport.Interceptor
ExtName string // Provider.Name(), captured at wrap time for *AbortError.Extension
}
// RoundTrip invokes the interceptor pre hook, restores the original context,
// executes the built-in chain (unless aborted), then calls the post hook if
// non-nil. When the extension implements AbortableInterceptor and returns a
// non-nil error from PreRoundTripE, the built-in chain is skipped and an
// *exttransport.AbortError is returned; the post hook is still invoked with
// (nil, reason) so extensions can unwind resources.
func (m *extensionMiddleware) RoundTrip(req *http.Request) (*http.Response, error) {
origCtx := req.Context()
req = req.Clone(origCtx)
var (
post func(*http.Response, error)
abortEr error
)
if a, ok := m.Ext.(exttransport.AbortableInterceptor); ok {
post, abortEr = a.PreRoundTripE(req)
} else {
post = m.Ext.PreRoundTrip(req)
}
if abortEr != nil {
if post != nil {
post(nil, abortEr)
}
return nil, &exttransport.AbortError{Extension: m.ExtName, Reason: abortEr}
}
req = req.WithContext(origCtx) // restore original context
resp, err := m.Base.RoundTrip(req)
if post != nil {
post(resp, err)
}
return resp, err
}
// wrapWithExtension wraps transport with the registered extension middleware.
// If no extension is registered, returns transport unchanged.
func wrapWithExtension(transport http.RoundTripper) http.RoundTripper {
p := exttransport.GetProvider()
if p == nil {
return transport
}
tr := p.ResolveInterceptor(context.Background())
if tr == nil {
return transport
}
return &extensionMiddleware{Base: transport, Ext: tr, ExtName: p.Name()}
}

View File

@@ -14,8 +14,8 @@ import (
"time"
exttransport "github.com/larksuite/cli/extension/transport"
internalauth "github.com/larksuite/cli/internal/auth"
"github.com/larksuite/cli/internal/riskcontrol"
internaltransport "github.com/larksuite/cli/internal/transport"
)
type roundTripFunc func(*http.Request) (*http.Response, error)
@@ -91,94 +91,107 @@ func TestRetryTransport_DefaultNoRetry(t *testing.T) {
}
}
// ---------------------------------------------------------------------------
// wrapSDKTransport chain composition
// buildSDKTransport policy behavior
// ---------------------------------------------------------------------------
func TestWrapSDKTransport_IncludesRetryTransport(t *testing.T) {
transport := wrapSDKTransport(riskcontrol.NewTransport(http.DefaultTransport, nil))
func TestBuildSDKTransportAppliesSecurityHeadersToEveryRequestClass(t *testing.T) {
exttransport.Register(nil)
received := make(chan http.Header, 2)
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, req *http.Request) {
received <- req.Header.Clone()
w.WriteHeader(http.StatusNoContent)
}))
t.Cleanup(server.Close)
// Chain: SecurityPolicy → BuildHeader → UserAgent → Retry → RiskControl → Base
sec, ok := transport.(*internalauth.SecurityPolicyTransport)
if !ok {
t.Fatalf("outer transport type = %T, want *auth.SecurityPolicyTransport", transport)
}
bh, ok := sec.Base.(*BuildHeaderTransport)
if !ok {
t.Fatalf("layer after SecurityPolicy = %T, want *BuildHeaderTransport", sec.Base)
}
ua, ok := bh.Base.(*UserAgentTransport)
if !ok {
t.Fatalf("layer after BuildHeader = %T, want *UserAgentTransport", bh.Base)
}
retry, ok := ua.Base.(*RetryTransport)
if !ok {
t.Fatalf("inner transport type = %T, want *RetryTransport", ua.Base)
}
if _, ok := retry.Base.(*riskcontrol.Transport); !ok {
t.Fatalf("layer after Retry = %T, want *riskcontrol.Transport", retry.Base)
for _, class := range []exttransport.RequestClass{
exttransport.RequestClassPlatform,
exttransport.RequestClassExternal,
} {
client := internaltransport.ClientForRequestClass(
&http.Client{Transport: buildSDKTransport(nil)},
class,
)
resp, err := client.Get(server.URL)
if err != nil {
t.Fatal(err)
}
resp.Body.Close()
headers := <-received
for header, values := range BaseSecurityHeaders() {
if len(values) == 0 {
continue
}
want := values[len(values)-1]
if got := headers.Get(header); got != want {
t.Fatalf("SDK %s header %s = %q, want %q", class, header, got, want)
}
}
}
}
func TestWrapSDKTransport_WithExtension(t *testing.T) {
func TestBuildSDKTransport_WithExtension(t *testing.T) {
previous := exttransport.GetProvider()
exttransport.Register(&stubTransportProvider{})
interceptor := &headerCapturingInterceptor{}
exttransport.Register(&platformOnlyStubProvider{
stubTransportProvider: &stubTransportProvider{interceptor: interceptor},
})
t.Cleanup(func() { exttransport.Register(previous) })
transport := wrapSDKTransport(riskcontrol.NewTransport(http.DefaultTransport, nil))
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, req *http.Request) {
w.WriteHeader(http.StatusNoContent)
}))
t.Cleanup(server.Close)
// Chain: extensionMiddleware → SecurityPolicy → BuildHeader → UserAgent → Retry → RiskControl → Base
mid, ok := transport.(*extensionMiddleware)
if !ok {
t.Fatalf("outer transport type = %T, want *extensionMiddleware", transport)
client := internaltransport.ClientForRequestClass(
&http.Client{Transport: buildSDKTransport(nil)},
exttransport.RequestClassPlatform,
)
resp, err := client.Get(server.URL)
if err != nil {
t.Fatal(err)
}
sec, ok := mid.Base.(*internalauth.SecurityPolicyTransport)
if !ok {
t.Fatalf("transport type = %T, want *auth.SecurityPolicyTransport", mid.Base)
}
bh, ok := sec.Base.(*BuildHeaderTransport)
if !ok {
t.Fatalf("layer after SecurityPolicy = %T, want *BuildHeaderTransport", sec.Base)
}
ua, ok := bh.Base.(*UserAgentTransport)
if !ok {
t.Fatalf("layer after BuildHeader = %T, want *UserAgentTransport", bh.Base)
}
retry, ok := ua.Base.(*RetryTransport)
if !ok {
t.Fatalf("innermost transport type = %T, want *RetryTransport", ua.Base)
}
if _, ok := retry.Base.(*riskcontrol.Transport); !ok {
t.Fatalf("layer after Retry = %T, want *riskcontrol.Transport", retry.Base)
resp.Body.Close()
if !interceptor.preCalled || !interceptor.postCalled {
t.Fatal("SDK platform request did not execute extension pre/post hooks")
}
}
func TestWrapSDKTransport_WithoutExtension(t *testing.T) {
func TestBuildSDKTransport_WithoutExtension(t *testing.T) {
previous := exttransport.GetProvider()
exttransport.Register(nil)
t.Cleanup(func() { exttransport.Register(previous) })
transport := wrapSDKTransport(riskcontrol.NewTransport(http.DefaultTransport, nil))
if _, ok := buildSDKTransport(nil).(*internaltransport.HTTPPolicyRouter); !ok {
t.Fatalf(
"buildSDKTransport() type = %T, want *transport.HTTPPolicyRouter",
buildSDKTransport(nil),
)
}
}
// Chain: SecurityPolicy → BuildHeader → UserAgent → Retry → RiskControl → Base
sec, ok := transport.(*internalauth.SecurityPolicyTransport)
func TestBuildSDKTransportSupportsPolicyLeafCloning(t *testing.T) {
previous := exttransport.GetProvider()
exttransport.Register(nil)
t.Cleanup(func() { exttransport.Register(previous) })
base := &http.Transport{}
client := internaltransport.ClientForRequestClass(
&http.Client{Transport: buildSDKTransportWithBase(base, nil)},
exttransport.RequestClassExternal,
)
source, ok := client.Transport.(interface {
CloneHTTPTransport() (http.RoundTripper, *http.Transport, bool)
})
if !ok {
t.Fatalf("outer transport type = %T, want *auth.SecurityPolicyTransport", transport)
t.Fatalf("SDK request-class transport type = %T, want clone capability", client.Transport)
}
bh, ok := sec.Base.(*BuildHeaderTransport)
if !ok {
t.Fatalf("layer after SecurityPolicy = %T, want *BuildHeaderTransport", sec.Base)
rebuilt, concrete, ok := source.CloneHTTPTransport()
if !ok || rebuilt == nil || concrete == nil {
t.Fatal("SDK policy graph could not clone its HTTP transport leaf")
}
ua, ok := bh.Base.(*UserAgentTransport)
if !ok {
t.Fatalf("layer after BuildHeader = %T, want *UserAgentTransport", bh.Base)
}
retry, ok := ua.Base.(*RetryTransport)
if !ok {
t.Fatalf("inner transport type = %T, want *RetryTransport", ua.Base)
}
if _, ok := retry.Base.(*riskcontrol.Transport); !ok {
t.Fatalf("layer after Retry = %T, want *riskcontrol.Transport", retry.Base)
if concrete == base {
t.Fatal("SDK policy graph reused the original HTTP transport")
}
}
@@ -238,7 +251,7 @@ func TestExtensionInterceptor_ExecutionOrder(t *testing.T) {
var base http.RoundTripper = http.DefaultTransport
base = &RetryTransport{Base: base}
base = &SecurityHeaderTransport{Base: base}
transport := wrapWithExtension(base)
transport := internaltransport.WrapWithExtension(base)
client := &http.Client{Transport: transport}
req, _ := http.NewRequest("GET", srv.URL, nil)
@@ -266,14 +279,16 @@ func TestExtensionInterceptor_ExecutionOrder(t *testing.T) {
}
}
// buildTamperingInterceptor tries to delete and spoof X-Cli-Build via
// PreRoundTrip. The SDK chain's BuildHeaderTransport must restore the real
// value before the request leaves the process.
// buildTamperingInterceptor tries to delete and spoof security headers via
// PreRoundTrip. The SDK built-in chain must restore the real values before the
// request leaves the process.
type buildTamperingInterceptor struct{}
func (buildTamperingInterceptor) PreRoundTrip(req *http.Request) func(*http.Response, error) {
req.Header.Del(HeaderBuild)
req.Header.Set(HeaderBuild, "ext-tampered-build")
req.Header.Del(HeaderSource)
req.Header.Set(HeaderSource, "ext-tampered-source")
return nil
}
@@ -285,7 +300,74 @@ func (riskHeaderTamperingInterceptor) PreRoundTrip(req *http.Request) func(*http
return nil
}
func TestWrapSDKTransport_StripsExtensionRiskHeaders(t *testing.T) {
type bootstrapPolicyTamperingInterceptor struct{}
func (bootstrapPolicyTamperingInterceptor) PreRoundTrip(req *http.Request) func(*http.Response, error) {
req.Header.Set(HeaderSource, "extension-value")
req.Header.Set(riskcontrol.HeaderOSType, "extension-value")
return nil
}
func TestNewDefaultInstallsSDKBootstrapSecurityPolicy(t *testing.T) {
t.Setenv("LARKSUITE_CLI_CONFIG_DIR", t.TempDir())
oldTransport := http.DefaultClient.Transport
oldCheckRedirect := http.DefaultClient.CheckRedirect
t.Cleanup(func() {
http.DefaultClient.Transport = oldTransport
http.DefaultClient.CheckRedirect = oldCheckRedirect
})
previous := exttransport.GetProvider()
exttransport.Register(&platformOnlyStubProvider{
stubTransportProvider: &stubTransportProvider{
interceptor: bootstrapPolicyTamperingInterceptor{},
},
})
t.Cleanup(func() { exttransport.Register(previous) })
var received http.Header
network := roundTripFunc(func(req *http.Request) (*http.Response, error) {
received = req.Header.Clone()
return &http.Response{
StatusCode: http.StatusNoContent,
Body: http.NoBody,
Request: req,
}, nil
})
http.DefaultClient.Transport = network
http.DefaultClient.CheckRedirect = nil
_ = NewDefault(nil, InvocationContext{})
req, err := http.NewRequest(
http.MethodPost,
"https://open.feishu.cn/callback/ws/endpoint",
strings.NewReader(`{"app_secret":"secret"}`),
)
if err != nil {
t.Fatal(err)
}
resp, err := http.DefaultClient.Do(req)
if err != nil {
t.Fatal(err)
}
resp.Body.Close()
if got := received.Get(HeaderSource); got != SourceValue {
t.Fatalf("%s = %q, want trusted value %q", HeaderSource, got, SourceValue)
}
if got := received.Get(riskcontrol.HeaderOSType); got != "" {
t.Fatalf("%s = %q, want extension value stripped", riskcontrol.HeaderOSType, got)
}
if got := received.Get(HeaderBuild); got != DetectBuildKind() {
t.Fatalf("%s = %q, want %q", HeaderBuild, got, DetectBuildKind())
}
if got := received.Get(HeaderUserAgent); got != UserAgentValue() {
t.Fatalf("%s = %q, want %q", HeaderUserAgent, got, UserAgentValue())
}
}
func TestBuildSDKTransport_StripsExtensionRiskHeaders(t *testing.T) {
previous := exttransport.GetProvider()
exttransport.Register(&stubTransportProvider{interceptor: riskHeaderTamperingInterceptor{}})
t.Cleanup(func() { exttransport.Register(previous) })
@@ -301,7 +383,11 @@ func TestWrapSDKTransport_StripsExtensionRiskHeaders(t *testing.T) {
}
req.Header.Set("Authorization", "Bearer token")
resp, err := wrapSDKTransport(riskcontrol.NewTransport(network, nil)).RoundTrip(req)
client := internaltransport.ClientForRequestClass(
&http.Client{Transport: buildSDKTransportWithBase(network, nil)},
exttransport.RequestClassPlatform,
)
resp, err := client.Do(req)
if err != nil {
t.Fatal(err)
}
@@ -312,14 +398,13 @@ func TestWrapSDKTransport_StripsExtensionRiskHeaders(t *testing.T) {
}
// TestBuildHeaderTransport_SDKChain_OverridesTamperedHeader verifies that the
// X-Cli-Build header is force-written by BuildHeaderTransport in the SDK
// transport chain, even when an extension tries to delete or spoof it. This
// closes the gap where the SDK chain had no equivalent of
// SecurityHeaderTransport (see design doc §3.3.3).
// SDK chain restores both the build classification and the full security
// header set after an extension runs.
func TestBuildHeaderTransport_SDKChain_OverridesTamperedHeader(t *testing.T) {
var receivedBuild string
var receivedBuild, receivedSource string
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
receivedBuild = r.Header.Get(HeaderBuild)
receivedSource = r.Header.Get(HeaderSource)
w.WriteHeader(http.StatusOK)
}))
defer srv.Close()
@@ -327,12 +412,13 @@ func TestBuildHeaderTransport_SDKChain_OverridesTamperedHeader(t *testing.T) {
exttransport.Register(&stubTransportProvider{interceptor: buildTamperingInterceptor{}})
t.Cleanup(func() { exttransport.Register(nil) })
// Replicate the SDK chain layering used by wrapSDKTransport.
// Replicate the SDK built-in chain inside buildSDKTransport.
var base http.RoundTripper = http.DefaultTransport
base = &RetryTransport{Base: base}
base = &UserAgentTransport{Base: base}
base = &BuildHeaderTransport{Base: base}
transport := wrapWithExtension(base)
base = &SecurityHeaderTransport{Base: base}
transport := internaltransport.WrapWithExtension(base)
client := &http.Client{Transport: transport}
req, _ := http.NewRequest("GET", srv.URL, nil)
@@ -349,6 +435,9 @@ func TestBuildHeaderTransport_SDKChain_OverridesTamperedHeader(t *testing.T) {
if receivedBuild != want {
t.Fatalf("%s = %q, want %q", HeaderBuild, receivedBuild, want)
}
if receivedSource != SourceValue {
t.Fatalf("%s = %q, want %q", HeaderSource, receivedSource, SourceValue)
}
}
// TestBuildHeaderTransport_OverridesEvenWithoutTamper verifies that even if
@@ -438,7 +527,7 @@ func TestExtensionInterceptor_ContextTamperPrevented(t *testing.T) {
return nil
})
mid := &extensionMiddleware{Base: capturer, Ext: tamperIC}
mid := &internaltransport.ExtensionMiddleware{Base: capturer, Ext: tamperIC}
origCtx := context.WithValue(context.Background(), testKey, "original")
req, _ := http.NewRequestWithContext(origCtx, "GET", srv.URL, nil)
@@ -500,7 +589,7 @@ func TestExtensionMiddleware_PreRoundTripEAbort(t *testing.T) {
return &http.Response{StatusCode: http.StatusOK, Body: http.NoBody}, nil
})
mid := &extensionMiddleware{Base: base, Ext: ic, ExtName: "stub"}
mid := &internaltransport.ExtensionMiddleware{Base: base, Ext: ic, ExtName: "stub"}
req, _ := http.NewRequest("GET", "http://example.invalid/", nil)
resp, err := mid.RoundTrip(req)
@@ -541,7 +630,7 @@ func TestExtensionMiddleware_PreRoundTripEAbort(t *testing.T) {
return nil, nil
})
mid := &extensionMiddleware{Base: base, Ext: ic, ExtName: "stub"}
mid := &internaltransport.ExtensionMiddleware{Base: base, Ext: ic, ExtName: "stub"}
req, _ := http.NewRequest("GET", "http://example.invalid/", nil)
_, err := mid.RoundTrip(req)
@@ -560,7 +649,7 @@ func TestExtensionMiddleware_PreRoundTripEHappyPath(t *testing.T) {
return &http.Response{StatusCode: http.StatusOK, Body: http.NoBody}, nil
})
mid := &extensionMiddleware{Base: base, Ext: ic, ExtName: "stub"}
mid := &internaltransport.ExtensionMiddleware{Base: base, Ext: ic, ExtName: "stub"}
req, _ := http.NewRequest("GET", "http://example.invalid/", nil)
resp, err := mid.RoundTrip(req)
if err != nil {

View File

@@ -3,7 +3,10 @@
package core
import "strings"
import (
"net/url"
"strings"
)
// LarkBrand represents the Lark platform brand.
// "feishu" targets China-mainland, "lark" targets international.
@@ -63,3 +66,39 @@ func ResolveEndpoints(brand LarkBrand) Endpoints {
func ResolveOpenBaseURL(brand LarkBrand) string {
return ResolveEndpoints(brand).Open
}
var platformEndpointHosts = func() map[string]struct{} {
hosts := make(map[string]struct{})
for _, brand := range []LarkBrand{BrandFeishu, BrandLark} {
endpoints := ResolveEndpoints(brand)
for _, rawURL := range []string{endpoints.Open, endpoints.Accounts, endpoints.MCP, endpoints.AppLink} {
parsed, err := url.Parse(rawURL)
if err == nil && parsed.Hostname() != "" {
hosts[strings.ToLower(parsed.Hostname())] = struct{}{}
}
}
}
return hosts
}()
// IsPlatformEndpointHost reports whether hostname exactly matches one of the
// endpoint hosts produced by ResolveEndpoints. It intentionally does not use a
// suffix match: lookalike external domains must never enter the platform
// transport extension.
func IsPlatformEndpointHost(hostname string) bool {
_, ok := platformEndpointHosts[strings.ToLower(hostname)]
return ok
}
// IsPlatformEndpointURL reports whether candidate uses a secure origin for a
// configured platform endpoint. Non-TLS and non-standard-port lookalikes are
// excluded even when their hostname matches.
func IsPlatformEndpointURL(candidate *url.URL) bool {
if candidate == nil || !strings.EqualFold(candidate.Scheme, "https") {
return false
}
if port := candidate.Port(); port != "" && port != "443" {
return false
}
return IsPlatformEndpointHost(candidate.Hostname())
}

View File

@@ -3,7 +3,11 @@
package core
import "testing"
import (
"net/url"
"reflect"
"testing"
)
func TestResolveEndpoints_Feishu(t *testing.T) {
ep := ResolveEndpoints(BrandFeishu)
@@ -91,3 +95,85 @@ func TestResolveEndpoints_NormalizesBrand(t *testing.T) {
t.Errorf("ResolveEndpoints(unexpected).Open = %q, want the feishu default", got)
}
}
func TestIsPlatformEndpointHost_ExactMatchOnly(t *testing.T) {
for _, host := range []string{
"open.feishu.cn",
"accounts.feishu.cn",
"mcp.feishu.cn",
"applink.feishu.cn",
"open.larksuite.com",
"accounts.larksuite.com",
"mcp.larksuite.com",
"applink.larksuite.com",
} {
if !IsPlatformEndpointHost(host) {
t.Errorf("IsPlatformEndpointHost(%q) = false, want true", host)
}
}
for _, host := range []string{
"example.com",
"open.feishu.cn.example.com",
"notopen.feishu.cn",
"",
} {
if IsPlatformEndpointHost(host) {
t.Errorf("IsPlatformEndpointHost(%q) = true, want false", host)
}
}
}
func TestIsPlatformEndpointHost_CoversEveryResolvedEndpoint(t *testing.T) {
for _, brand := range []LarkBrand{BrandFeishu, BrandLark} {
endpoints := reflect.ValueOf(ResolveEndpoints(brand))
for i := 0; i < endpoints.NumField(); i++ {
rawURL := endpoints.Field(i).String()
parsed, err := url.Parse(rawURL)
if err != nil {
t.Fatalf("ResolveEndpoints(%q) field %d URL %q: %v", brand, i, rawURL, err)
}
if !IsPlatformEndpointHost(parsed.Hostname()) {
t.Errorf("ResolveEndpoints(%q) field %d host %q is missing from the platform transport boundary", brand, i, parsed.Hostname())
}
}
}
}
func TestIsPlatformEndpointURL_RequiresSecureStandardOrigin(t *testing.T) {
if IsPlatformEndpointURL(nil) {
t.Error("IsPlatformEndpointURL(nil) = true, want false")
}
uppercaseScheme := &url.URL{Scheme: "HTTPS", Host: "open.feishu.cn", Path: "/path"}
if !IsPlatformEndpointURL(uppercaseScheme) {
t.Error("IsPlatformEndpointURL() rejected uppercase HTTPS scheme")
}
for _, rawURL := range []string{
"http://open.feishu.cn/path",
"https://open.feishu.cn:8443/path",
"https://open.feishu.cn.example.com/path",
} {
candidate, err := url.Parse(rawURL)
if err != nil {
t.Fatal(err)
}
if IsPlatformEndpointURL(candidate) {
t.Errorf("IsPlatformEndpointURL(%q) = true, want false", rawURL)
}
}
for _, rawURL := range []string{
"https://open.feishu.cn/path",
"https://open.feishu.cn:443/path",
"https://OPEN.FEISHU.CN/path",
} {
candidate, err := url.Parse(rawURL)
if err != nil {
t.Fatal(err)
}
if !IsPlatformEndpointURL(candidate) {
t.Errorf("IsPlatformEndpointURL(%q) = false, want true", rawURL)
}
}
}

View File

@@ -180,8 +180,8 @@ func saveCachedMerged(data []byte, cm CacheMeta) error {
// localVersion is sent as data_version query param for server-side version comparison.
// Returns (data, reg, err). A nil reg means the version is unchanged (not modified).
func fetchRemoteMerged(localVersion string) (data []byte, reg *MergedRegistry, err error) {
// Route through the shared proxy-plugin-aware transport so remote API
// definition fetches honor proxy plugin mode instead of bypassing it.
// Remote metadata is platform traffic and must honor both the shared proxy
// configuration and the registered platform transport extension.
client := transport.NewHTTPClient(fetchTimeout)
req, err := http.NewRequest("GET", remoteMetaURL(localVersion), nil)
if err != nil {

View File

@@ -12,6 +12,8 @@ import (
internaltransport "github.com/larksuite/cli/internal/transport"
)
var _ internaltransport.RoundTripperDecorator = (*Transport)(nil)
const (
HeaderProductModel = "X-Agent-Device-Type"
HeaderOSType = "X-Agent-Os-Type"
@@ -40,6 +42,28 @@ func NewTransport(next http.RoundTripper, source Source) *Transport {
}
}
// BaseRoundTripper exposes the network transport so policy routers can clone
// and rebuild the complete decorator graph without dropping risk control.
func (t *Transport) BaseRoundTripper() http.RoundTripper {
if t == nil || t.next == nil {
return internaltransport.Fallback()
}
return t.next
}
// WithBaseRoundTripper returns an equivalent risk-control boundary over base.
func (t *Transport) WithBaseRoundTripper(base http.RoundTripper) http.RoundTripper {
if t == nil {
return NewTransport(base, nil)
}
cloned := *t
if base == nil {
base = internaltransport.Fallback()
}
cloned.next = base
return &cloned
}
// RoundTrip implements http.RoundTripper.
func (t *Transport) RoundTrip(req *http.Request) (*http.Response, error) {
req = req.Clone(req.Context())

View File

@@ -2,7 +2,7 @@
// SPDX-License-Identifier: MIT
// Package transport owns how the CLI assembles its outbound HTTP transport: the
// shared base RoundTripper (Shared/Fallback/NewHTTPClient), the LARK_CLI_NO_PROXY
// shared base RoundTripper (Shared/Fallback and the HTTP client constructors), the LARK_CLI_NO_PROXY
// direct-egress clone, and the ~/.lark-cli/proxy_config.json proxy-plugin mode.
//
// Proxy-plugin mode forces all outbound HTTP(S) requests through a fixed loopback

View File

@@ -0,0 +1,258 @@
// Copyright (c) 2026 Lark Technologies Pte. Ltd.
// SPDX-License-Identifier: MIT
package transport
import (
"context"
"net/http"
"net/url"
"strings"
"sync"
larkws "github.com/larksuite/oapi-sdk-go/v3/ws"
"github.com/larksuite/cli/errs"
exttransport "github.com/larksuite/cli/extension/transport"
"github.com/larksuite/cli/internal/core"
)
type requestMatcher func(*http.Request) bool
type transportPolicyBuilder func(http.RoundTripper) http.RoundTripper
type sdkBootstrapRedirectContextKey struct{}
var (
// larkws pins this client during package initialization.
sdkBootstrapHTTPClient = http.DefaultClient
installDefaultClientMu sync.Mutex
)
// sdkBootstrapTransport applies the platform HTTP policy only to dependency
// bootstrap requests selected by match. Unmatched DefaultClient traffic is
// delegated directly to the previous transport.
type sdkBootstrapTransport struct {
base http.RoundTripper
match requestMatcher
buildPlatformPolicy transportPolicyBuilder
policyMu sync.RWMutex
}
func (t *sdkBootstrapTransport) RoundTrip(req *http.Request) (*http.Response, error) {
if !t.isBootstrapRequest(req) {
return t.fallbackTransport().RoundTrip(req)
}
base := t.base
if base == nil {
// Resolve Shared lazily so bridge installation never initializes
// workspace-scoped proxy state ahead of workspace selection.
base = Shared()
}
buildPlatformPolicy := t.platformPolicyBuilder()
if buildPlatformPolicy == nil {
return nil, errs.NewInternalError(
errs.SubtypeUnknown,
"SDK bootstrap transport policy is not configured",
)
}
base = buildPlatformPolicy(base)
if base == nil {
return nil, errs.NewInternalError(
errs.SubtypeUnknown,
"SDK bootstrap transport policy returned a nil transport",
)
}
// Resolve extensions per hop so redirects retain platform policy.
extended := WrapWithExtensionForClass(base, exttransport.RequestClassPlatform)
guarded := &sameOriginRedirectTransport{base: extended}
return guarded.RoundTrip(req)
}
func (t *sdkBootstrapTransport) platformPolicyBuilder() transportPolicyBuilder {
t.policyMu.RLock()
defer t.policyMu.RUnlock()
return t.buildPlatformPolicy
}
func (t *sdkBootstrapTransport) setPlatformPolicyBuilder(build transportPolicyBuilder) {
t.policyMu.Lock()
t.buildPlatformPolicy = build
t.policyMu.Unlock()
}
func (t *sdkBootstrapTransport) isBootstrapRequest(req *http.Request) bool {
if req == nil {
return false
}
if _, redirected := req.Context().Value(sdkBootstrapRedirectContextKey{}).(struct{}); redirected {
return true
}
return t.match != nil && t.match(req)
}
func (t *sdkBootstrapTransport) fallbackTransport() http.RoundTripper {
if t.base != nil {
return t.base
}
// Preserve net/http's dynamic nil-Transport fallback.
return http.DefaultTransport
}
// sameOriginRedirectTransport rejects redirects before net/http can replay a
// bootstrap request to a different logical origin.
type sameOriginRedirectTransport struct {
base http.RoundTripper
}
func (t *sameOriginRedirectTransport) RoundTrip(req *http.Request) (*http.Response, error) {
resp, err := t.base.RoundTrip(req)
if err != nil || resp == nil || !isFollowedRedirect(resp.StatusCode) {
return resp, err
}
location := resp.Header.Get("Location")
if location == "" {
return resp, nil
}
target, parseErr := req.URL.Parse(location)
if parseErr != nil {
if resp.Body != nil {
_ = resp.Body.Close()
}
return nil, errs.NewInternalError(
errs.SubtypeInvalidResponse,
"platform request returned an invalid redirect location: %v",
parseErr,
).WithCause(parseErr)
}
if sameOrigin(req.URL, target) {
return resp, nil
}
if resp.Body != nil {
_ = resp.Body.Close()
}
return nil, errs.NewSecurityPolicyError(
errs.SubtypeAccessDenied,
"platform bootstrap blocked cross-origin redirect from %q to %q",
originName(req.URL),
originName(target),
)
}
// sdkBootstrapRedirectPolicy preserves the prior hook and marks each redirect hop.
func sdkBootstrapRedirectPolicy(
match requestMatcher,
previous func(*http.Request, []*http.Request) error,
) func(*http.Request, []*http.Request) error {
return func(req *http.Request, via []*http.Request) error {
if previous != nil {
if err := previous(req, via); err != nil {
return err
}
} else if len(via) >= 10 {
// Retain net/http's default redirect limit.
return errs.NewNetworkError(
errs.SubtypeNetworkTransport,
"stopped after 10 redirects",
)
}
if req == nil || len(via) == 0 || match == nil || !match(via[0]) {
return nil
}
ctx := context.WithValue(req.Context(), sdkBootstrapRedirectContextKey{}, struct{}{})
*req = *req.WithContext(ctx)
return nil
}
}
func originName(candidate *url.URL) string {
if candidate == nil {
return ""
}
return strings.ToLower(candidate.Scheme) + "://" + candidate.Host
}
func sameOrigin(left, right *url.URL) bool {
if left == nil || right == nil {
return false
}
return strings.EqualFold(left.Scheme, right.Scheme) &&
strings.EqualFold(left.Hostname(), right.Hostname()) &&
originPort(left) == originPort(right)
}
func originPort(candidate *url.URL) string {
if port := candidate.Port(); port != "" {
return port
}
switch strings.ToLower(candidate.Scheme) {
case "http":
return "80"
case "https":
return "443"
default:
return ""
}
}
func isFollowedRedirect(status int) bool {
switch status {
case http.StatusMovedPermanently,
http.StatusFound,
http.StatusSeeOther,
http.StatusTemporaryRedirect,
http.StatusPermanentRedirect:
return true
default:
return false
}
}
// InstallSDKTransportBridge wraps larkws's captured HTTP bootstrap client. All
// requests through that client hit the bridge, but only matched bootstrap
// traffic uses platform policy. The SDK owns the subsequent WebSocket dial,
// which does not use this net/http transport.
func InstallSDKTransportBridge(buildPlatformPolicy func(http.RoundTripper) http.RoundTripper) {
installDefaultClientMu.Lock()
defer installDefaultClientMu.Unlock()
installSDKTransportBridge(
sdkBootstrapHTTPClient,
isSDKWebSocketBootstrapRequest,
buildPlatformPolicy,
)
}
func isSDKWebSocketBootstrapRequest(req *http.Request) bool {
return req != nil &&
req.Method == http.MethodPost &&
core.IsPlatformEndpointURL(req.URL) &&
req.URL.Path == larkws.GenEndpointUri
}
func installSDKTransportBridge(
client *http.Client,
match requestMatcher,
buildPlatformPolicy transportPolicyBuilder,
) {
if client == nil {
return
}
if existing, ok := client.Transport.(*sdkBootstrapTransport); ok {
existing.setPlatformPolicyBuilder(buildPlatformPolicy)
return
}
base := client.Transport
previousRedirect := client.CheckRedirect
client.Transport = &sdkBootstrapTransport{
base: base,
match: match,
buildPlatformPolicy: buildPlatformPolicy,
}
client.CheckRedirect = sdkBootstrapRedirectPolicy(match, previousRedirect)
}

View File

@@ -0,0 +1,120 @@
// Copyright (c) 2026 Lark Technologies Pte. Ltd.
// SPDX-License-Identifier: MIT
package transport
import (
"context"
"net/http"
exttransport "github.com/larksuite/cli/extension/transport"
)
var _ RoundTripperDecorator = (*ExtensionMiddleware)(nil)
type resolvedExtension struct {
provider exttransport.Provider
interceptor exttransport.Interceptor
}
func resolveExtension() *resolvedExtension {
p := exttransport.GetProvider()
if p == nil {
return nil
}
interceptor := p.ResolveInterceptor(context.Background())
if interceptor == nil {
return nil
}
return &resolvedExtension{provider: p, interceptor: interceptor}
}
func (e *resolvedExtension) wrap(base http.RoundTripper, class exttransport.RequestClass, enforceScope bool) http.RoundTripper {
if base == nil {
base = Shared()
}
if e == nil {
return base
}
if enforceScope {
if scoped, ok := e.provider.(exttransport.ScopedProvider); ok && !scoped.SupportsRequestClass(class) {
return base
}
}
return &ExtensionMiddleware{Base: base, Ext: e.interceptor, ExtName: e.provider.Name()}
}
// ExtensionMiddleware wraps the built-in transport chain with extension
// pre/post hooks. The built-in chain always executes unless an
// exttransport.AbortableInterceptor rejects the request.
//
// The original request context is restored after the pre hook to prevent an
// extension from replacing cancellation, deadlines, or built-in values. The
// request is cloned so URL and header mutations do not alter the caller's
// request object. The body remains shared; interceptors that consume it must
// restore it before returning.
type ExtensionMiddleware struct {
Base http.RoundTripper
Ext exttransport.Interceptor
ExtName string
}
// BaseRoundTripper returns the wrapped built-in transport chain.
func (m *ExtensionMiddleware) BaseRoundTripper() http.RoundTripper {
if m.Base == nil {
return Shared()
}
return m.Base
}
// WithBaseRoundTripper clones the middleware over base.
func (m *ExtensionMiddleware) WithBaseRoundTripper(base http.RoundTripper) http.RoundTripper {
cloned := *m
cloned.Base = base
return &cloned
}
// RoundTrip invokes the extension pre hook, the wrapped transport, and then
// the optional post hook. Abortable interceptors can stop the request before
// the wrapped transport is called.
func (m *ExtensionMiddleware) RoundTrip(req *http.Request) (*http.Response, error) {
origCtx := req.Context()
req = req.Clone(origCtx)
var (
post func(*http.Response, error)
abortErr error
)
if a, ok := m.Ext.(exttransport.AbortableInterceptor); ok {
post, abortErr = a.PreRoundTripE(req)
} else {
post = m.Ext.PreRoundTrip(req)
}
if abortErr != nil {
if post != nil {
post(nil, abortErr)
}
return nil, &exttransport.AbortError{Extension: m.ExtName, Reason: abortErr}
}
req = req.WithContext(origCtx)
resp, err := m.BaseRoundTripper().RoundTrip(req)
if post != nil {
post(resp, err)
}
return resp, err
}
// WrapWithExtension wraps base with the currently registered transport
// extension. With no registered provider or no resolved interceptor, base is
// returned unchanged.
func WrapWithExtension(base http.RoundTripper) http.RoundTripper {
return resolveExtension().wrap(base, "", false)
}
// WrapWithExtensionForClass wraps base only when the registered provider
// supports class. Providers without the optional ScopedProvider interface keep
// their historical all-request behavior.
func WrapWithExtensionForClass(base http.RoundTripper, class exttransport.RequestClass) http.RoundTripper {
return resolveExtension().wrap(base, class, true)
}

View File

@@ -0,0 +1,924 @@
// Copyright (c) 2026 Lark Technologies Pte. Ltd.
// SPDX-License-Identifier: MIT
package transport
import (
"context"
"errors"
"io"
"net/http"
"net/http/httptest"
"net/url"
"strconv"
"strings"
"sync/atomic"
"testing"
larkws "github.com/larksuite/oapi-sdk-go/v3/ws"
"github.com/larksuite/cli/errs"
exttransport "github.com/larksuite/cli/extension/transport"
)
type testProvider struct {
interceptor exttransport.Interceptor
resolveCalls *int
}
func (p testProvider) Name() string { return "test-provider" }
func (p testProvider) ResolveInterceptor(context.Context) exttransport.Interceptor {
if p.resolveCalls != nil {
*p.resolveCalls++
}
return p.interceptor
}
type scopedTestProvider struct {
testProvider
supported exttransport.RequestClass
}
func (p scopedTestProvider) SupportsRequestClass(class exttransport.RequestClass) bool {
return class == p.supported
}
type testHeaderInterceptor struct {
calls int
}
func (i *testHeaderInterceptor) PreRoundTrip(req *http.Request) func(*http.Response, error) {
i.calls++
req.Header.Set("X-Test-Platform", "routed")
return nil
}
type abortingTestInterceptor struct {
reason error
post func(*http.Response, error)
}
func (i *abortingTestInterceptor) PreRoundTrip(*http.Request) func(*http.Response, error) {
panic("PreRoundTrip called for abortable interceptor")
}
func (i *abortingTestInterceptor) PreRoundTripE(*http.Request) (func(*http.Response, error), error) {
return i.post, i.reason
}
func TestLegacyProviderKeepsAllRequestBehavior(t *testing.T) {
t.Setenv("LARKSUITE_CLI_CONFIG_DIR", t.TempDir())
unsetProxyPluginEnv(t)
resetProxyPluginState()
t.Setenv(EnvNoProxy, "")
interceptor := &testHeaderInterceptor{}
previousProvider := exttransport.GetProvider()
exttransport.Register(testProvider{interceptor: interceptor})
t.Cleanup(func() { exttransport.Register(previousProvider) })
received := make(chan string, 2)
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, req *http.Request) {
received <- req.Header.Get("X-Test-Platform")
w.WriteHeader(http.StatusNoContent)
}))
t.Cleanup(server.Close)
for _, client := range []*http.Client{
ClientForRequestClass(NewHTTPClient(0), exttransport.RequestClassPlatform),
NewExternalHTTPClient(0),
} {
resp, err := client.Get(server.URL)
if err != nil {
t.Fatal(err)
}
resp.Body.Close()
}
if got := <-received; got != "routed" {
t.Fatalf("platform request header = %q, want routed", got)
}
if got := <-received; got != "routed" {
t.Fatalf("external request header = %q, want routed for legacy provider", got)
}
if interceptor.calls != 2 {
t.Fatalf("extension calls = %d, want exactly 2", interceptor.calls)
}
}
func TestScopedProviderOnlyRunsForSupportedRequestClass(t *testing.T) {
t.Setenv("LARKSUITE_CLI_CONFIG_DIR", t.TempDir())
unsetProxyPluginEnv(t)
resetProxyPluginState()
t.Setenv(EnvNoProxy, "")
interceptor := &testHeaderInterceptor{}
previousProvider := exttransport.GetProvider()
exttransport.Register(scopedTestProvider{
testProvider: testProvider{interceptor: interceptor},
supported: exttransport.RequestClassPlatform,
})
t.Cleanup(func() { exttransport.Register(previousProvider) })
received := make(chan string, 2)
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, req *http.Request) {
received <- req.Header.Get("X-Test-Platform")
w.WriteHeader(http.StatusNoContent)
}))
t.Cleanup(server.Close)
clients := []*http.Client{
ClientForRequestClass(NewHTTPClient(0), exttransport.RequestClassPlatform),
NewExternalHTTPClient(0),
}
for _, client := range clients {
resp, err := client.Get(server.URL)
if err != nil {
t.Fatal(err)
}
resp.Body.Close()
}
if got := <-received; got != "routed" {
t.Fatalf("platform request header = %q, want routed", got)
}
if got := <-received; got != "" {
t.Fatalf("external request received scoped provider header %q", got)
}
if interceptor.calls != 1 {
t.Fatalf("extension calls = %d, want exactly 1", interceptor.calls)
}
}
func TestHTTPPolicyRouterResolvesProviderOnce(t *testing.T) {
resolveCalls := 0
previousProvider := exttransport.GetProvider()
exttransport.Register(testProvider{
interceptor: &testHeaderInterceptor{},
resolveCalls: &resolveCalls,
})
t.Cleanup(func() { exttransport.Register(previousProvider) })
base := roundTripFunc(func(req *http.Request) (*http.Response, error) {
return &http.Response{StatusCode: http.StatusNoContent, Body: http.NoBody, Request: req}, nil
})
_ = NewHTTPPolicyRouter(base, base)
if resolveCalls != 1 {
t.Fatalf("ResolveInterceptor() calls = %d, want 1 per router", resolveCalls)
}
}
func TestSDKBootstrapBridgeBlocksCrossOriginRedirectAfterSameOriginHop(t *testing.T) {
var externalCalls atomic.Int32
var relayBody string
base := roundTripFunc(func(req *http.Request) (*http.Response, error) {
if req.URL.Host == "external.example" {
externalCalls.Add(1)
return noContentResponse(req), nil
}
switch req.URL.Path {
case "/bootstrap":
return redirectResponse(req, http.StatusTemporaryRedirect, "/relay"), nil
case "/relay":
body, err := io.ReadAll(req.Body)
if err != nil {
return nil, err
}
relayBody = string(body)
return redirectResponse(
req,
http.StatusPermanentRedirect,
"https://external.example/target",
), nil
default:
return noContentResponse(req), nil
}
})
client := &http.Client{Transport: base}
installSDKTransportBridge(client, func(req *http.Request) bool {
return req.URL != nil && req.URL.Path == "/bootstrap"
}, identityTransportPolicy)
const secret = "app_secret=secret"
req, err := http.NewRequest(
http.MethodPost,
"https://platform.example/bootstrap",
strings.NewReader(secret),
)
if err != nil {
t.Fatal(err)
}
resp, err := client.Do(req)
if resp != nil && resp.Body != nil {
resp.Body.Close()
}
if err == nil || !strings.Contains(err.Error(), "cross-origin redirect") {
t.Fatalf("Do() error = %v, want cross-origin redirect rejection", err)
}
if problem, ok := errs.ProblemOf(err); !ok ||
problem.Category != errs.CategoryPolicy ||
problem.Subtype != errs.SubtypeAccessDenied {
t.Fatalf("Do() problem = %#v, %v; want policy/access_denied", problem, ok)
}
if relayBody != secret {
t.Fatalf("same-origin relay body = %q, want %q", relayBody, secret)
}
if got := externalCalls.Load(); got != 0 {
t.Fatalf("cross-origin target calls = %d, want 0", got)
}
}
func TestSDKBootstrapRedirectGuardClassifiesInvalidLocation(t *testing.T) {
base := roundTripFunc(func(req *http.Request) (*http.Response, error) {
return redirectResponse(req, http.StatusFound, "%"), nil
})
client := &http.Client{Transport: &sameOriginRedirectTransport{base: base}}
resp, err := client.Get("https://platform.example/bootstrap")
if resp != nil && resp.Body != nil {
resp.Body.Close()
}
if err == nil || !strings.Contains(err.Error(), "invalid redirect location") {
t.Fatalf("Do() error = %v, want invalid redirect rejection", err)
}
problem, ok := errs.ProblemOf(err)
if !ok || problem.Category != errs.CategoryInternal || problem.Subtype != errs.SubtypeInvalidResponse {
t.Fatalf("Do() problem = %#v, %v; want internal/invalid_response", problem, ok)
}
}
type redirectPolicyInterceptor struct {
calls int
}
func (i *redirectPolicyInterceptor) PreRoundTrip(req *http.Request) func(*http.Response, error) {
i.calls++
req.Header.Set("X-Extension-Hop", strconv.Itoa(i.calls))
req.Header.Set("X-Reserved", "extension")
return nil
}
func TestSDKBootstrapBridgeRetainsPoliciesAcrossSameOriginRedirect(t *testing.T) {
previousProvider := exttransport.GetProvider()
interceptor := &redirectPolicyInterceptor{}
exttransport.Register(scopedTestProvider{
testProvider: testProvider{interceptor: interceptor},
supported: exttransport.RequestClassPlatform,
})
t.Cleanup(func() { exttransport.Register(previousProvider) })
var finalHeaders http.Header
base := roundTripFunc(func(req *http.Request) (*http.Response, error) {
switch req.URL.Path {
case "/bootstrap":
return redirectResponse(req, http.StatusTemporaryRedirect, "/next"), nil
case "/next":
finalHeaders = req.Header.Clone()
return noContentResponse(req), nil
default:
return noContentResponse(req), nil
}
})
builtInCalls := 0
client := &http.Client{Transport: base}
installSDKTransportBridge(
client,
func(req *http.Request) bool {
return req.URL != nil && req.URL.Path == "/bootstrap"
},
func(base http.RoundTripper) http.RoundTripper {
return roundTripFunc(func(req *http.Request) (*http.Response, error) {
builtInCalls++
req = req.Clone(req.Context())
req.Header.Set("X-Builtin-Hop", strconv.Itoa(builtInCalls))
req.Header.Set("X-Reserved", "trusted")
return base.RoundTrip(req)
})
},
)
req, err := http.NewRequest(
http.MethodPost,
"https://platform.example/bootstrap",
strings.NewReader("body"),
)
if err != nil {
t.Fatal(err)
}
resp, err := client.Do(req)
if err != nil {
t.Fatal(err)
}
resp.Body.Close()
if finalHeaders == nil {
t.Fatal("same-origin redirect target was not called")
}
if interceptor.calls != 2 {
t.Fatalf("extension calls = %d, want 2", interceptor.calls)
}
if builtInCalls != 2 {
t.Fatalf("built-in policy calls = %d, want 2", builtInCalls)
}
if got := finalHeaders.Get("X-Extension-Hop"); got != "2" {
t.Fatalf("final X-Extension-Hop = %q, want 2", got)
}
if got := finalHeaders.Get("X-Builtin-Hop"); got != "2" {
t.Fatalf("final X-Builtin-Hop = %q, want 2", got)
}
if got := finalHeaders.Get("X-Reserved"); got != "trusted" {
t.Fatalf("final X-Reserved = %q, want trusted built-in value", got)
}
}
type redirectRewriteInterceptor struct {
target *url.URL
postLocation string
calls int
}
func (i *redirectRewriteInterceptor) PreRoundTrip(req *http.Request) func(*http.Response, error) {
i.calls++
req.URL.Scheme = i.target.Scheme
req.URL.Host = i.target.Host
if i.postLocation == "" {
return nil
}
return func(resp *http.Response, err error) {
if err == nil && resp != nil && isFollowedRedirect(resp.StatusCode) {
resp.Header.Set("Location", i.postLocation)
}
}
}
func TestSDKBootstrapRedirectGuardUsesLogicalURLAfterExtensionRewrite(t *testing.T) {
sidecarURL, err := url.Parse("https://sidecar.example")
if err != nil {
t.Fatal(err)
}
sidecarCalls := 0
base := roundTripFunc(func(req *http.Request) (*http.Response, error) {
if req.URL.Host != sidecarURL.Host {
t.Fatalf("network host = %q, want extension target %q", req.URL.Host, sidecarURL.Host)
}
sidecarCalls++
switch req.URL.Path {
case "/bootstrap":
return redirectResponse(
req,
http.StatusTemporaryRedirect,
"https://platform.example/next",
), nil
case "/next":
return noContentResponse(req), nil
default:
return noContentResponse(req), nil
}
})
previousProvider := exttransport.GetProvider()
interceptor := &redirectRewriteInterceptor{target: sidecarURL}
exttransport.Register(scopedTestProvider{
testProvider: testProvider{interceptor: interceptor},
supported: exttransport.RequestClassPlatform,
})
t.Cleanup(func() { exttransport.Register(previousProvider) })
client := &http.Client{Transport: base}
installSDKTransportBridge(client, func(req *http.Request) bool {
return req.URL != nil && req.URL.Path == "/bootstrap"
}, identityTransportPolicy)
req, err := http.NewRequest(
http.MethodPost,
"https://platform.example/bootstrap",
strings.NewReader("body"),
)
if err != nil {
t.Fatal(err)
}
resp, err := client.Do(req)
if err != nil {
t.Fatal(err)
}
resp.Body.Close()
if sidecarCalls != 2 {
t.Fatalf("sidecar calls = %d, want 2", sidecarCalls)
}
if interceptor.calls != 2 {
t.Fatalf("extension calls = %d, want 2", interceptor.calls)
}
}
func TestSDKBootstrapRedirectGuardChecksLocationAfterExtensionPostHook(t *testing.T) {
var externalCalls atomic.Int32
sidecarURL, err := url.Parse("https://sidecar.example")
if err != nil {
t.Fatal(err)
}
base := roundTripFunc(func(req *http.Request) (*http.Response, error) {
if req.URL.Host == "external.example" {
externalCalls.Add(1)
return noContentResponse(req), nil
}
return redirectResponse(
req,
http.StatusTemporaryRedirect,
"https://platform.example/next",
), nil
})
previousProvider := exttransport.GetProvider()
exttransport.Register(scopedTestProvider{
testProvider: testProvider{interceptor: &redirectRewriteInterceptor{
target: sidecarURL,
postLocation: "https://external.example/target",
}},
supported: exttransport.RequestClassPlatform,
})
t.Cleanup(func() { exttransport.Register(previousProvider) })
client := &http.Client{Transport: base}
installSDKTransportBridge(client, func(req *http.Request) bool {
return req.URL != nil && req.URL.Path == "/bootstrap"
}, identityTransportPolicy)
req, err := http.NewRequest(
http.MethodPost,
"https://platform.example/bootstrap",
strings.NewReader("secret"),
)
if err != nil {
t.Fatal(err)
}
resp, err := client.Do(req)
if resp != nil && resp.Body != nil {
resp.Body.Close()
}
if err == nil || !strings.Contains(err.Error(), "cross-origin redirect") {
t.Fatalf("Do() error = %v, want post-hook Location rejection", err)
}
if problem, ok := errs.ProblemOf(err); !ok ||
problem.Category != errs.CategoryPolicy ||
problem.Subtype != errs.SubtypeAccessDenied {
t.Fatalf("Do() problem = %#v, %v; want policy/access_denied", problem, ok)
}
if got := externalCalls.Load(); got != 0 {
t.Fatalf("post-hook redirect target calls = %d, want 0", got)
}
}
func TestSameOriginNormalizesDefaultPort(t *testing.T) {
left, err := url.Parse("https://platform.example/bootstrap")
if err != nil {
t.Fatal(err)
}
right, err := url.Parse("https://platform.example:443/next")
if err != nil {
t.Fatal(err)
}
if !sameOrigin(left, right) {
t.Fatal("sameOrigin() = false for equivalent default HTTPS ports")
}
}
func TestDefaultClientBridgeCoversWebSocketSDKBootstrap(t *testing.T) {
preserveHTTPClientState(t, sdkBootstrapHTTPClient)
t.Setenv("LARKSUITE_CLI_CONFIG_DIR", t.TempDir())
unsetProxyPluginEnv(t)
resetProxyPluginState()
t.Setenv(EnvNoProxy, "1")
previousProvider := exttransport.GetProvider()
interceptor := &testHeaderInterceptor{}
exttransport.Register(scopedTestProvider{
testProvider: testProvider{interceptor: interceptor},
supported: exttransport.RequestClassPlatform,
})
t.Cleanup(func() { exttransport.Register(previousProvider) })
seenHeader := make(chan string, 1)
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, req *http.Request) {
seenHeader <- req.Header.Get("X-Test-Platform")
w.Header().Set("Content-Type", "application/json")
w.WriteHeader(http.StatusBadRequest)
_, _ = io.WriteString(w, `{"code":400,"msg":"stop after bootstrap"}`)
}))
t.Cleanup(server.Close)
installSDKTransportBridge(sdkBootstrapHTTPClient, func(req *http.Request) bool {
return req.URL != nil && req.URL.Host == strings.TrimPrefix(server.URL, "http://")
}, identityTransportPolicy)
client := larkws.NewClient(
"test-app",
"test-secret",
larkws.WithDomain(server.URL),
larkws.WithAutoReconnect(false),
)
if err := client.Start(context.Background()); err == nil {
t.Fatal("WebSocket SDK Start() error = nil, want bootstrap failure")
}
if got := <-seenHeader; got != "routed" {
t.Fatalf("WebSocket bootstrap header = %q, want routed", got)
}
if interceptor.calls != 1 {
t.Fatalf("extension calls = %d, want exactly 1 bootstrap call", interceptor.calls)
}
}
func TestSDKTransportBridgeUsesPinnedClientAfterGlobalReplacement(t *testing.T) {
preserveHTTPClientState(t, sdkBootstrapHTTPClient)
oldDefaultClient := http.DefaultClient
t.Cleanup(func() { http.DefaultClient = oldDefaultClient })
var pinnedCalls atomic.Int32
pinnedHeader := make(chan string, 1)
sdkBootstrapHTTPClient.Transport = roundTripFunc(func(req *http.Request) (*http.Response, error) {
pinnedCalls.Add(1)
pinnedHeader <- req.Header.Get("X-Pinned-Bridge")
return &http.Response{
StatusCode: http.StatusBadRequest,
Header: http.Header{"Content-Type": []string{"application/json"}},
Body: io.NopCloser(strings.NewReader(`{"code":400,"msg":"stop"}`)),
Request: req,
}, nil
})
sdkBootstrapHTTPClient.CheckRedirect = nil
var replacementCalls atomic.Int32
http.DefaultClient = &http.Client{Transport: roundTripFunc(func(req *http.Request) (*http.Response, error) {
replacementCalls.Add(1)
return &http.Response{
StatusCode: http.StatusBadRequest,
Body: http.NoBody,
Request: req,
}, nil
})}
InstallSDKTransportBridge(func(base http.RoundTripper) http.RoundTripper {
return roundTripFunc(func(req *http.Request) (*http.Response, error) {
req = req.Clone(req.Context())
req.Header.Set("X-Pinned-Bridge", "routed")
return base.RoundTrip(req)
})
})
client := larkws.NewClient(
"test-app",
"test-secret",
larkws.WithAutoReconnect(false),
)
if err := client.Start(context.Background()); err == nil {
t.Fatal("WebSocket SDK Start() error = nil, want bootstrap failure")
}
if got := pinnedCalls.Load(); got != 1 {
t.Fatalf("SDK-pinned client calls = %d, want 1", got)
}
if got := <-pinnedHeader; got != "routed" {
t.Fatalf("SDK-pinned bridge header = %q, want routed", got)
}
if got := replacementCalls.Load(); got != 0 {
t.Fatalf("replacement DefaultClient calls = %d, want 0", got)
}
}
func TestSDKWebSocketBootstrapMatcherIsNarrow(t *testing.T) {
tests := []struct {
name string
method string
url string
want bool
}{
{
name: "platform bootstrap",
method: http.MethodPost,
url: "https://open.feishu.cn/callback/ws/endpoint",
want: true,
},
{
name: "other platform path",
method: http.MethodPost,
url: "https://open.feishu.cn/open-apis/test",
},
{
name: "wrong bootstrap method",
method: http.MethodGet,
url: "https://open.feishu.cn/callback/ws/endpoint",
},
{
name: "external lookalike",
method: http.MethodPost,
url: "https://external.example/callback/ws/endpoint",
},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
req, err := http.NewRequest(tt.method, tt.url, nil)
if err != nil {
t.Fatal(err)
}
if got := isSDKWebSocketBootstrapRequest(req); got != tt.want {
t.Fatalf("isSDKWebSocketBootstrapRequest() = %v, want %v", got, tt.want)
}
})
}
}
func TestSDKTransportBridgeLeavesOtherPlatformPathsUntouched(t *testing.T) {
previousProvider := exttransport.GetProvider()
interceptor := &testHeaderInterceptor{}
exttransport.Register(scopedTestProvider{
testProvider: testProvider{interceptor: interceptor},
supported: exttransport.RequestClassPlatform,
})
t.Cleanup(func() { exttransport.Register(previousProvider) })
baseCalls := 0
client := &http.Client{Transport: roundTripFunc(func(req *http.Request) (*http.Response, error) {
baseCalls++
return &http.Response{StatusCode: http.StatusNoContent, Body: http.NoBody, Request: req}, nil
})}
installSDKTransportBridge(client, isSDKWebSocketBootstrapRequest, nil)
req, err := http.NewRequest(http.MethodGet, "https://open.feishu.cn/open-apis/test", nil)
if err != nil {
t.Fatal(err)
}
resp, err := client.Do(req)
if err != nil {
t.Fatal(err)
}
resp.Body.Close()
if baseCalls != 1 {
t.Fatalf("base calls = %d, want 1", baseCalls)
}
if interceptor.calls != 0 {
t.Fatalf("extension calls = %d, want 0 for unmatched DefaultClient traffic", interceptor.calls)
}
}
func TestSDKTransportBridgeNilBasePreservesDefaultTransportForUnmatchedRequest(t *testing.T) {
oldDefaultTransport := http.DefaultTransport
t.Cleanup(func() { http.DefaultTransport = oldDefaultTransport })
unsetProxyPluginEnv(t)
resetProxyPluginState()
t.Setenv(EnvNoProxy, "1")
var firstCalls atomic.Int32
http.DefaultTransport = roundTripFunc(func(req *http.Request) (*http.Response, error) {
firstCalls.Add(1)
return &http.Response{StatusCode: http.StatusNoContent, Body: http.NoBody, Request: req}, nil
})
client := &http.Client{}
installSDKTransportBridge(client, func(*http.Request) bool { return false }, nil)
var currentCalls atomic.Int32
http.DefaultTransport = roundTripFunc(func(req *http.Request) (*http.Response, error) {
currentCalls.Add(1)
return &http.Response{StatusCode: http.StatusNoContent, Body: http.NoBody, Request: req}, nil
})
req, err := http.NewRequest(http.MethodGet, "http://127.0.0.1:1/unmatched", nil)
if err != nil {
t.Fatal(err)
}
resp, err := client.Do(req)
if err != nil {
t.Fatal(err)
}
resp.Body.Close()
if got := firstCalls.Load(); got != 0 {
t.Fatalf("install-time DefaultTransport calls = %d, want 0", got)
}
if got := currentCalls.Load(); got != 1 {
t.Fatalf("request-time DefaultTransport calls = %d, want 1", got)
}
}
func TestSDKTransportBridgeUpdatesPlatformPolicy(t *testing.T) {
client := &http.Client{Transport: roundTripFunc(func(req *http.Request) (*http.Response, error) {
return noContentResponse(req), nil
})}
var firstCalls, secondCalls int
build := func(calls *int) transportPolicyBuilder {
return func(base http.RoundTripper) http.RoundTripper {
*calls++
return base
}
}
match := func(*http.Request) bool { return true }
installSDKTransportBridge(client, match, build(&firstCalls))
installSDKTransportBridge(client, match, build(&secondCalls))
req, err := http.NewRequest(http.MethodPost, "https://platform.example/bootstrap", nil)
if err != nil {
t.Fatal(err)
}
resp, err := client.Do(req)
if err != nil {
t.Fatal(err)
}
resp.Body.Close()
if firstCalls != 0 || secondCalls != 1 {
t.Fatalf("policy calls = (%d, %d), want (0, 1)", firstCalls, secondCalls)
}
}
func TestSDKBootstrapTransportFailsClosedWithoutPlatformPolicy(t *testing.T) {
var baseCalls atomic.Int32
client := &http.Client{Transport: roundTripFunc(func(req *http.Request) (*http.Response, error) {
baseCalls.Add(1)
return &http.Response{
StatusCode: http.StatusNoContent,
Body: http.NoBody,
Request: req,
}, nil
})}
installSDKTransportBridge(client, func(*http.Request) bool { return true }, nil)
req, err := http.NewRequest(http.MethodPost, "https://platform.example/bootstrap", nil)
if err != nil {
t.Fatal(err)
}
resp, err := client.Do(req)
if resp != nil && resp.Body != nil {
resp.Body.Close()
}
if err == nil || !strings.Contains(err.Error(), "policy is not configured") {
t.Fatalf("Do() error = %v, want missing policy rejection", err)
}
if problem, ok := errs.ProblemOf(err); !ok ||
problem.Category != errs.CategoryInternal ||
problem.Subtype != errs.SubtypeUnknown {
t.Fatalf("Do() problem = %#v, %v; want internal/unknown", problem, ok)
}
if got := baseCalls.Load(); got != 0 {
t.Fatalf("base transport calls = %d, want 0", got)
}
}
func TestSDKBootstrapTransportFailsClosedForNilPlatformTransport(t *testing.T) {
var baseCalls atomic.Int32
client := &http.Client{Transport: roundTripFunc(func(req *http.Request) (*http.Response, error) {
baseCalls.Add(1)
return &http.Response{
StatusCode: http.StatusNoContent,
Body: http.NoBody,
Request: req,
}, nil
})}
installSDKTransportBridge(client, func(*http.Request) bool { return true }, func(http.RoundTripper) http.RoundTripper {
return nil
})
req, err := http.NewRequest(http.MethodPost, "https://platform.example/bootstrap", nil)
if err != nil {
t.Fatal(err)
}
resp, err := client.Do(req)
if resp != nil && resp.Body != nil {
resp.Body.Close()
}
if err == nil || !strings.Contains(err.Error(), "nil transport") {
t.Fatalf("Do() error = %v, want nil policy transport rejection", err)
}
if problem, ok := errs.ProblemOf(err); !ok ||
problem.Category != errs.CategoryInternal ||
problem.Subtype != errs.SubtypeUnknown {
t.Fatalf("Do() problem = %#v, %v; want internal/unknown", problem, ok)
}
if got := baseCalls.Load(); got != 0 {
t.Fatalf("base transport calls = %d, want 0", got)
}
}
func TestSDKBootstrapRedirectPolicyRetainsDefaultLimit(t *testing.T) {
policy := sdkBootstrapRedirectPolicy(nil, nil)
via := make([]*http.Request, 10)
err := policy(&http.Request{}, via)
if err == nil {
t.Fatal("redirect policy error = nil after 10 redirects")
}
if problem, ok := errs.ProblemOf(err); !ok ||
problem.Category != errs.CategoryNetwork ||
problem.Subtype != errs.SubtypeNetworkTransport {
t.Fatalf("redirect problem = %#v, %v; want network/transport", problem, ok)
}
}
func TestExtensionMiddlewareUsesFallbackWhenBaseIsNil(t *testing.T) {
t.Setenv("LARKSUITE_CLI_CONFIG_DIR", t.TempDir())
unsetProxyPluginEnv(t)
resetProxyPluginState()
t.Setenv(EnvNoProxy, "")
previous := http.DefaultTransport
var calls atomic.Int32
http.DefaultTransport = roundTripFunc(func(req *http.Request) (*http.Response, error) {
calls.Add(1)
return &http.Response{
StatusCode: http.StatusNoContent,
Body: http.NoBody,
Request: req,
}, nil
})
t.Cleanup(func() { http.DefaultTransport = previous })
req, err := http.NewRequest(http.MethodGet, "https://external.example/file", nil)
if err != nil {
t.Fatal(err)
}
resp, err := (&ExtensionMiddleware{Ext: &testHeaderInterceptor{}}).RoundTrip(req)
if err != nil {
t.Fatal(err)
}
resp.Body.Close()
if got := calls.Load(); got != 1 {
t.Fatalf("fallback transport calls = %d, want 1", got)
}
}
func TestExtensionMiddlewareAbortsBeforeBase(t *testing.T) {
reason := errors.New("blocked")
baseCalled := false
postCalled := false
interceptor := &abortingTestInterceptor{
reason: reason,
post: func(resp *http.Response, err error) {
postCalled = true
if resp != nil || err != reason {
t.Errorf("post arguments = (%v, %v), want (nil, reason)", resp, err)
}
},
}
middleware := &ExtensionMiddleware{
Base: roundTripFunc(func(*http.Request) (*http.Response, error) {
baseCalled = true
return nil, nil
}),
Ext: interceptor,
ExtName: "test-provider",
}
resp, err := middleware.RoundTrip(httptest.NewRequest(http.MethodGet, "https://example.com", nil))
if resp != nil {
t.Fatalf("response = %v, want nil", resp)
}
var abortErr *exttransport.AbortError
if !errors.As(err, &abortErr) {
t.Fatalf("error = %T, want *transport.AbortError", err)
}
if abortErr.Extension != "test-provider" || abortErr.Reason != reason {
t.Fatalf("abort error = %#v, want provider and reason", abortErr)
}
if baseCalled {
t.Fatal("base transport was called")
}
if !postCalled {
t.Fatal("post hook was not called")
}
}
func preserveHTTPClientState(t *testing.T, client *http.Client) {
t.Helper()
oldTransport := client.Transport
oldCheckRedirect := client.CheckRedirect
t.Cleanup(func() {
client.Transport = oldTransport
client.CheckRedirect = oldCheckRedirect
})
}
func identityTransportPolicy(base http.RoundTripper) http.RoundTripper {
return base
}
func redirectResponse(req *http.Request, status int, location string) *http.Response {
return &http.Response{
StatusCode: status,
Header: http.Header{"Location": []string{location}},
Body: http.NoBody,
Request: req,
}
}
func noContentResponse(req *http.Request) *http.Response {
return &http.Response{
StatusCode: http.StatusNoContent,
Body: http.NoBody,
Request: req,
}
}
type roundTripFunc func(*http.Request) (*http.Response, error)
func (f roundTripFunc) RoundTrip(req *http.Request) (*http.Response, error) {
return f(req)
}

View File

@@ -0,0 +1,232 @@
// Copyright (c) 2026 Lark Technologies Pte. Ltd.
// SPDX-License-Identifier: MIT
package transport
import (
"context"
"net/http"
"github.com/larksuite/cli/errs"
exttransport "github.com/larksuite/cli/extension/transport"
"github.com/larksuite/cli/internal/core"
)
type requestClassContextKey struct{}
type forcedRequestClassContextKey struct{}
// HTTPPolicyRouter selects an HTTP transport policy from request intent and
// the endpoint catalog. Explicit request intent takes precedence; otherwise
// known platform endpoints use the platform policy and all other URLs use the
// external policy.
type HTTPPolicyRouter struct {
platform http.RoundTripper
external http.RoundTripper
}
// RoundTripperDecorator describes a transport layer that can be rebuilt over
// a cloned base transport. Connection-policy helpers use this contract to
// preserve retry, response, and extension layers while safely customizing the
// innermost *http.Transport.
type RoundTripperDecorator interface {
BaseRoundTripper() http.RoundTripper
WithBaseRoundTripper(http.RoundTripper) http.RoundTripper
}
// NewHTTPPolicyRouter constructs a router over two policy chains. A nil chain
// falls back to the shared proxy-aware transport. The currently registered
// extension provider is resolved once and applied according to its optional
// ScopedProvider contract.
func NewHTTPPolicyRouter(platform, external http.RoundTripper) *HTTPPolicyRouter {
if platform == nil {
platform = Shared()
}
if external == nil {
external = Shared()
}
extension := resolveExtension()
return &HTTPPolicyRouter{
platform: extension.wrap(platform, exttransport.RequestClassPlatform, true),
external: extension.wrap(external, exttransport.RequestClassExternal, true),
}
}
// RoundTrip dispatches the request to its selected policy chain.
func (r *HTTPPolicyRouter) RoundTrip(req *http.Request) (*http.Response, error) {
if req == nil {
return nil, errs.NewInternalError(
errs.SubtypeUnknown,
"HTTP policy router received a nil request",
)
}
class, err := classifyRequest(req)
if err != nil {
return nil, err
}
if class == exttransport.RequestClassPlatform {
return r.platform.RoundTrip(req)
}
return r.external.RoundTrip(req)
}
func (r *HTTPPolicyRouter) transportForClass(class exttransport.RequestClass) (http.RoundTripper, bool) {
switch class {
case exttransport.RequestClassPlatform:
return r.platform, true
case exttransport.RequestClassExternal:
return r.external, true
default:
return nil, false
}
}
func classifyRequest(req *http.Request) (exttransport.RequestClass, error) {
if explicit, ok := req.Context().Value(requestClassContextKey{}).(exttransport.RequestClass); ok {
switch explicit {
case exttransport.RequestClassPlatform, exttransport.RequestClassExternal:
return explicit, nil
default:
return "", errs.NewInternalError(
errs.SubtypeUnknown,
"unsupported HTTP request class %q",
explicit,
)
}
}
if core.IsPlatformEndpointURL(req.URL) {
return exttransport.RequestClassPlatform, nil
}
return exttransport.RequestClassExternal, nil
}
// WithRequestClass returns a shallow copy of req with explicit routing intent.
func WithRequestClass(req *http.Request, class exttransport.RequestClass) *http.Request {
if req == nil {
return nil
}
ctx := context.WithValue(req.Context(), requestClassContextKey{}, class)
return req.WithContext(ctx)
}
func withForcedRequestClass(req *http.Request, class exttransport.RequestClass) *http.Request {
if req == nil {
return nil
}
if _, forced := req.Context().Value(forcedRequestClassContextKey{}).(struct{}); forced {
return req
}
ctx := context.WithValue(req.Context(), requestClassContextKey{}, class)
ctx = context.WithValue(ctx, forcedRequestClassContextKey{}, struct{}{})
return req.WithContext(ctx)
}
type requestClassTransport struct {
base http.RoundTripper
class exttransport.RequestClass
}
func (t *requestClassTransport) RoundTrip(req *http.Request) (*http.Response, error) {
return t.base.RoundTrip(withForcedRequestClass(req, t.class))
}
// CloneHTTPTransport exposes a structural cloning capability without requiring
// higher-level safety helpers to import this package. The explicit request
// class selects the policy branch that must be rebuilt.
func (t *requestClassTransport) CloneHTTPTransport() (http.RoundTripper, *http.Transport, bool) {
return CloneHTTPTransportForRequestClass(t.base, t.class)
}
// TransformHTTPTransport clones the selected policy branch and replaces its
// concrete transport in place. Keeping the replacement at the graph leaf is
// important for policies that must observe requests after outer decorators
// have run, such as proxy selection.
func (t *requestClassTransport) TransformHTTPTransport(transform func(*http.Transport) (http.RoundTripper, bool)) (http.RoundTripper, bool) {
return transformHTTPTransportForRequestClass(t.base, t.class, transform, 0)
}
// ClientForRequestClass clones client and forces all of its requests through a
// specific policy class. The original client is never mutated.
func ClientForRequestClass(client *http.Client, class exttransport.RequestClass) *http.Client {
if client == nil {
client = &http.Client{}
}
cloned := *client
base := client.Transport
if base == nil {
base = Shared()
}
cloned.Transport = &requestClassTransport{base: base, class: class}
return &cloned
}
// CloneHTTPTransportForRequestClass selects one policy branch, clones its
// innermost *http.Transport, and rebuilds every composable decorator around
// the clone. Callers can customize concrete before using rebuilt. The original
// transport graph is never mutated.
func CloneHTTPTransportForRequestClass(base http.RoundTripper, class exttransport.RequestClass) (rebuilt http.RoundTripper, concrete *http.Transport, ok bool) {
rebuilt, ok = transformHTTPTransportForRequestClass(base, class, func(cloned *http.Transport) (http.RoundTripper, bool) {
concrete = cloned
return cloned, true
}, 0)
if !ok {
return nil, nil, false
}
return rebuilt, concrete, true
}
func transformHTTPTransportForRequestClass(
base http.RoundTripper,
class exttransport.RequestClass,
transform func(*http.Transport) (http.RoundTripper, bool),
depth int,
) (http.RoundTripper, bool) {
if depth > 32 {
return nil, false
}
if base == nil || transform == nil {
if transform == nil {
return nil, false
}
base = Shared()
}
switch current := base.(type) {
case *http.Transport:
cloned := cloneHTTPTransport(current)
rebuilt, valid := transform(cloned)
return rebuilt, valid && rebuilt != nil
case *requestClassTransport:
return transformHTTPTransportForRequestClass(current.base, class, transform, depth+1)
case *HTTPPolicyRouter:
selected, valid := current.transportForClass(class)
if !valid {
return nil, false
}
return transformHTTPTransportForRequestClass(selected, class, transform, depth+1)
case RoundTripperDecorator:
inner := current.BaseRoundTripper()
if inner == nil || inner == base {
return nil, false
}
rebuiltInner, valid := transformHTTPTransportForRequestClass(inner, class, transform, depth+1)
if !valid {
return nil, false
}
rebuilt := current.WithBaseRoundTripper(rebuiltInner)
return rebuilt, rebuilt != nil
default:
return nil, false
}
}
func cloneHTTPTransport(source *http.Transport) *http.Transport {
cloned := source.Clone()
// Clone leaves an auto-configured h2 handler on source.
if cloned.TLSNextProto == nil {
if _, ok := source.TLSNextProto["h2"]; ok {
cloned.ForceAttemptHTTP2 = true
}
}
return cloned
}

View File

@@ -0,0 +1,351 @@
// Copyright (c) 2026 Lark Technologies Pte. Ltd.
// SPDX-License-Identifier: MIT
package transport
import (
"errors"
"net/http"
"net/http/httptest"
"net/url"
"strings"
"testing"
"github.com/larksuite/cli/errs"
exttransport "github.com/larksuite/cli/extension/transport"
)
type cloneTestDecorator struct {
base http.RoundTripper
}
func (d *cloneTestDecorator) RoundTrip(req *http.Request) (*http.Response, error) {
return d.base.RoundTrip(req)
}
func (d *cloneTestDecorator) BaseRoundTripper() http.RoundTripper {
return d.base
}
func (d *cloneTestDecorator) WithBaseRoundTripper(base http.RoundTripper) http.RoundTripper {
return &cloneTestDecorator{base: base}
}
type headerCloneTestDecorator struct {
base http.RoundTripper
}
func (d *headerCloneTestDecorator) RoundTrip(req *http.Request) (*http.Response, error) {
req = req.Clone(req.Context())
req.Header.Set("X-Decorator", "applied")
return d.base.RoundTrip(req)
}
func (d *headerCloneTestDecorator) BaseRoundTripper() http.RoundTripper {
return d.base
}
func (d *headerCloneTestDecorator) WithBaseRoundTripper(base http.RoundTripper) http.RoundTripper {
return &headerCloneTestDecorator{base: base}
}
func TestHTTPPolicyRouterClassifiesFromEndpointCatalog(t *testing.T) {
exttransport.Register(nil)
platformCalls := 0
externalCalls := 0
router := NewHTTPPolicyRouter(
roundTripFunc(func(req *http.Request) (*http.Response, error) {
platformCalls++
return &http.Response{StatusCode: http.StatusNoContent, Body: http.NoBody, Request: req}, nil
}),
roundTripFunc(func(req *http.Request) (*http.Response, error) {
externalCalls++
return &http.Response{StatusCode: http.StatusNoContent, Body: http.NoBody, Request: req}, nil
}),
)
for _, rawURL := range []string{
"https://open.feishu.cn/open-apis/test",
"https://example.com/file",
} {
req, err := http.NewRequest(http.MethodGet, rawURL, nil)
if err != nil {
t.Fatal(err)
}
resp, err := router.RoundTrip(req)
if err != nil {
t.Fatal(err)
}
resp.Body.Close()
}
if platformCalls != 1 || externalCalls != 1 {
t.Fatalf("platform calls = %d, external calls = %d; want 1 each", platformCalls, externalCalls)
}
}
func TestHTTPPolicyRouterExplicitClassOverridesCatalog(t *testing.T) {
exttransport.Register(nil)
platformCalls := 0
externalCalls := 0
router := NewHTTPPolicyRouter(
roundTripFunc(func(req *http.Request) (*http.Response, error) {
platformCalls++
return &http.Response{StatusCode: http.StatusNoContent, Body: http.NoBody, Request: req}, nil
}),
roundTripFunc(func(req *http.Request) (*http.Response, error) {
externalCalls++
return &http.Response{StatusCode: http.StatusNoContent, Body: http.NoBody, Request: req}, nil
}),
)
req, err := http.NewRequest(http.MethodGet, "https://open.feishu.cn/open-apis/test", nil)
if err != nil {
t.Fatal(err)
}
req = WithRequestClass(req, exttransport.RequestClassExternal)
resp, err := router.RoundTrip(req)
if err != nil {
t.Fatal(err)
}
resp.Body.Close()
if platformCalls != 0 || externalCalls != 1 {
t.Fatalf("platform calls = %d, external calls = %d; want 0 and 1", platformCalls, externalCalls)
}
}
func TestClientForRequestClassOutermostIntentWins(t *testing.T) {
exttransport.Register(nil)
platformCalls := 0
externalCalls := 0
router := NewHTTPPolicyRouter(
roundTripFunc(func(req *http.Request) (*http.Response, error) {
platformCalls++
return &http.Response{StatusCode: http.StatusNoContent, Body: http.NoBody, Request: req}, nil
}),
roundTripFunc(func(req *http.Request) (*http.Response, error) {
externalCalls++
return &http.Response{StatusCode: http.StatusNoContent, Body: http.NoBody, Request: req}, nil
}),
)
platform := ClientForRequestClass(&http.Client{Transport: router}, exttransport.RequestClassPlatform)
external := ClientForRequestClass(platform, exttransport.RequestClassExternal)
resp, err := external.Get("https://open.feishu.cn/open-apis/test")
if err != nil {
t.Fatal(err)
}
resp.Body.Close()
if platformCalls != 0 || externalCalls != 1 {
t.Fatalf("platform calls = %d, external calls = %d; want outer external intent to win", platformCalls, externalCalls)
}
}
func TestHTTPPolicyRouterRejectsInvalidExplicitClass(t *testing.T) {
exttransport.Register(nil)
router := NewHTTPPolicyRouter(nil, nil)
req, err := http.NewRequest(http.MethodGet, "https://example.com", nil)
if err != nil {
t.Fatal(err)
}
req = WithRequestClass(req, exttransport.RequestClass("invalid"))
if _, err := router.RoundTrip(req); err == nil || !strings.Contains(err.Error(), "unsupported HTTP request class") {
t.Fatalf("RoundTrip() error = %v, want unsupported request class", err)
} else if problem, ok := errs.ProblemOf(err); !ok ||
problem.Category != errs.CategoryInternal ||
problem.Subtype != errs.SubtypeUnknown {
t.Fatalf("RoundTrip() problem = %#v, %v; want internal/unknown", problem, ok)
}
}
func TestHTTPPolicyRouterRejectsNilRequest(t *testing.T) {
router := NewHTTPPolicyRouter(nil, nil)
_, err := router.RoundTrip(nil)
problem, ok := errs.ProblemOf(err)
if !ok || problem.Category != errs.CategoryInternal || problem.Subtype != errs.SubtypeUnknown {
t.Fatalf("RoundTrip() problem = %#v, %v; want internal/unknown", problem, ok)
}
}
func TestHTTPPolicyRouterReclassifiesRedirectTargets(t *testing.T) {
interceptor := &testHeaderInterceptor{}
exttransport.Register(scopedTestProvider{
testProvider: testProvider{interceptor: interceptor},
supported: exttransport.RequestClassPlatform,
})
t.Cleanup(func() { exttransport.Register(nil) })
receivedHeader := make(chan string, 1)
external := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, req *http.Request) {
receivedHeader <- req.Header.Get("X-Test-Platform")
w.WriteHeader(http.StatusNoContent)
}))
t.Cleanup(external.Close)
router := NewHTTPPolicyRouter(
roundTripFunc(func(req *http.Request) (*http.Response, error) {
return &http.Response{
StatusCode: http.StatusFound,
Header: http.Header{"Location": []string{external.URL}},
Body: http.NoBody,
Request: req,
}, nil
}),
http.DefaultTransport,
)
client := &http.Client{Transport: router}
req, err := http.NewRequest(http.MethodGet, "https://open.feishu.cn/start", nil)
if err != nil {
t.Fatal(err)
}
resp, err := client.Do(req)
if err != nil {
t.Fatal(err)
}
resp.Body.Close()
if got := <-receivedHeader; got != "" {
t.Fatalf("redirect target received platform-scoped header %q", got)
}
if interceptor.calls != 1 {
t.Fatalf("extension calls = %d, want only the initial platform request", interceptor.calls)
}
}
func TestCloneHTTPTransportForRequestClassRebuildsDecorators(t *testing.T) {
wantErr := errors.New("preserved proxy policy")
base := &http.Transport{
Proxy: func(*http.Request) (*url.URL, error) {
return nil, wantErr
},
}
decorated := &cloneTestDecorator{base: base}
router := NewHTTPPolicyRouter(decorated, decorated)
rebuilt, concrete, ok := CloneHTTPTransportForRequestClass(router, exttransport.RequestClassExternal)
if !ok {
t.Fatal("CloneHTTPTransportForRequestClass() ok = false")
}
if concrete == base {
t.Fatal("CloneHTTPTransportForRequestClass() reused the original *http.Transport")
}
if _, ok := rebuilt.(*cloneTestDecorator); !ok {
t.Fatalf("rebuilt transport type = %T, want *cloneTestDecorator", rebuilt)
}
req, err := http.NewRequest(http.MethodGet, "https://external.example/file", nil)
if err != nil {
t.Fatal(err)
}
if _, err := rebuilt.RoundTrip(req); !errors.Is(err, wantErr) {
t.Fatalf("RoundTrip() error = %v, want %v", err, wantErr)
}
}
func TestCloneHTTPTransportForRequestClassPreservesAutomaticHTTP2(t *testing.T) {
previousProvider := exttransport.GetProvider()
exttransport.Register(nil)
t.Cleanup(func() { exttransport.Register(previousProvider) })
source := &http.Transport{
Proxy: http.ProxyURL(&url.URL{Scheme: "http", Host: "proxy.example:8080"}),
}
router := NewHTTPPolicyRouter(&http.Transport{}, source)
_, cloned, ok := CloneHTTPTransportForRequestClass(router, exttransport.RequestClassExternal)
if !ok {
t.Fatal("CloneHTTPTransportForRequestClass() ok = false")
}
if !cloned.ForceAttemptHTTP2 {
t.Fatal("ForceAttemptHTTP2 = false, want true")
}
if cloned.TLSNextProto != nil {
t.Fatal("TLSNextProto is non-nil, want automatic HTTP/2")
}
}
func TestCloneHTTPTransportForRequestClassKeepsOutermostIntent(t *testing.T) {
platformErr := errors.New("platform transport")
externalErr := errors.New("external transport")
newBlocked := func(reason error) *http.Transport {
return &http.Transport{Proxy: func(*http.Request) (*url.URL, error) { return nil, reason }}
}
router := NewHTTPPolicyRouter(newBlocked(platformErr), newBlocked(externalErr))
platform := ClientForRequestClass(&http.Client{Transport: router}, exttransport.RequestClassPlatform)
external := ClientForRequestClass(platform, exttransport.RequestClassExternal)
source, ok := external.Transport.(interface {
CloneHTTPTransport() (http.RoundTripper, *http.Transport, bool)
})
if !ok {
t.Fatalf("transport type %T has no clone capability", external.Transport)
}
rebuilt, _, ok := source.CloneHTTPTransport()
if !ok {
t.Fatal("CloneHTTPTransport() ok = false")
}
req, err := http.NewRequest(http.MethodGet, "https://open.feishu.cn/open-apis/test", nil)
if err != nil {
t.Fatal(err)
}
if _, err := rebuilt.RoundTrip(req); !errors.Is(err, externalErr) {
t.Fatalf("RoundTrip() error = %v, want outer external transport error %v", err, externalErr)
}
}
func TestClientForRequestClassOverridesCallerIntent(t *testing.T) {
platformErr := errors.New("platform transport")
externalErr := errors.New("external transport")
newBlocked := func(reason error) *http.Transport {
return &http.Transport{Proxy: func(*http.Request) (*url.URL, error) { return nil, reason }}
}
router := NewHTTPPolicyRouter(newBlocked(platformErr), newBlocked(externalErr))
client := ClientForRequestClass(&http.Client{Transport: router}, exttransport.RequestClassExternal)
req, err := http.NewRequest(http.MethodGet, "https://external.example/file", nil)
if err != nil {
t.Fatal(err)
}
req = WithRequestClass(req, exttransport.RequestClassPlatform)
if _, err := client.Do(req); !errors.Is(err, externalErr) {
t.Fatalf("Do() error = %v, want forced external transport error %v", err, externalErr)
}
}
func TestTransformHTTPTransportReplacesLeafInsideDecorators(t *testing.T) {
exttransport.Register(nil)
decorated := &headerCloneTestDecorator{base: &http.Transport{}}
router := NewHTTPPolicyRouter(decorated, decorated)
client := ClientForRequestClass(&http.Client{Transport: router}, exttransport.RequestClassExternal)
source, ok := client.Transport.(interface {
TransformHTTPTransport(func(*http.Transport) (http.RoundTripper, bool)) (http.RoundTripper, bool)
})
if !ok {
t.Fatalf("transport type %T has no transform capability", client.Transport)
}
rebuilt, ok := source.TransformHTTPTransport(func(*http.Transport) (http.RoundTripper, bool) {
return roundTripFunc(func(req *http.Request) (*http.Response, error) {
if got := req.Header.Get("X-Decorator"); got != "applied" {
t.Fatalf("leaf received X-Decorator = %q, want applied", got)
}
return &http.Response{StatusCode: http.StatusNoContent, Body: http.NoBody, Request: req}, nil
}), true
})
if !ok {
t.Fatal("TransformHTTPTransport() ok = false")
}
req, err := http.NewRequest(http.MethodGet, "https://external.example/file", nil)
if err != nil {
t.Fatal(err)
}
resp, err := rebuilt.RoundTrip(req)
if err != nil {
t.Fatal(err)
}
resp.Body.Close()
}

View File

@@ -8,6 +8,8 @@ import (
"os"
"sync"
"time"
exttransport "github.com/larksuite/cli/extension/transport"
)
// Shared returns the base http.RoundTripper for all CLI HTTP clients.
@@ -55,21 +57,29 @@ func Fallback() *http.Transport {
return noProxyTransport()
}
// NewHTTPClient returns an *http.Client whose Transport is the shared,
// proxy-plugin-aware base (see Shared). Prefer this over a bare &http.Client{}
// for outbound requests: a bare client falls back to http.DefaultTransport and
// therefore silently bypasses proxy plugin mode (fixed proxy + trusted CA, or
// fail-closed), creating an audit blind spot.
// NewHTTPClient returns a policy-routed client over the shared proxy-aware
// transport. Known platform endpoints use the platform request class; all
// other URLs use the external request class. Existing unscoped transport
// providers continue to apply to both classes.
//
// A zero timeout means no client-level timeout (callers relying on context
// deadlines pass 0).
func NewHTTPClient(timeout time.Duration) *http.Client {
base := Shared()
return &http.Client{
Transport: Shared(),
Transport: NewHTTPPolicyRouter(base, base),
Timeout: timeout,
}
}
// NewExternalHTTPClient returns a client for user-provided, pre-signed, CDN,
// package-registry, and other non-platform URLs. It forces the external policy
// while preserving the shared proxy configuration and the historical behavior
// of unscoped transport providers. A zero timeout means no client-level timeout.
func NewExternalHTTPClient(timeout time.Duration) *http.Client {
return ClientForRequestClass(NewHTTPClient(timeout), exttransport.RequestClassExternal)
}
// noProxyTransport is a proxy-disabled clone of http.DefaultTransport, lazily
// built the first time LARK_CLI_NO_PROXY is observed set.
var noProxyTransport = sync.OnceValue(func() *http.Transport {

View File

@@ -88,23 +88,24 @@ func TestShared_NoProxyOverridesSystemProxy(t *testing.T) {
}
}
// TestNewHTTPClient verifies the factory wires the shared proxy-plugin-aware
// transport (instead of a bare client that bypasses proxy plugin mode).
func TestNewHTTPClient(t *testing.T) {
// TestHTTPClientConstructors verifies both the policy-routed client and its
// forced-external view retain explicit transports and configured timeouts.
func TestHTTPClientConstructors(t *testing.T) {
t.Setenv("LARKSUITE_CLI_CONFIG_DIR", t.TempDir())
unsetProxyPluginEnv(t)
resetProxyPluginState()
t.Setenv(EnvNoProxy, "")
c := NewHTTPClient(7 * time.Second)
if c.Transport == nil {
t.Fatal("NewHTTPClient transport is nil; want shared transport")
}
if c.Transport != Shared() {
t.Errorf("NewHTTPClient transport = %v, want Shared()", c.Transport)
}
if c.Timeout != 7*time.Second {
t.Errorf("NewHTTPClient timeout = %v, want 7s", c.Timeout)
for name, client := range map[string]*http.Client{
"routed": NewHTTPClient(7 * time.Second),
"external": NewExternalHTTPClient(7 * time.Second),
} {
if client.Transport == nil {
t.Fatalf("%s client transport is nil", name)
}
if client.Timeout != 7*time.Second {
t.Errorf("%s client timeout = %v, want 7s", name, client.Timeout)
}
}
}
@@ -153,4 +154,32 @@ func TestShared_MalformedConfigFailsClosedEvenWithNoProxy(t *testing.T) {
if err == nil {
t.Fatalf("RoundTrip() err = nil (resp=%v); malformed config must fail closed", resp)
}
for name, test := range map[string]struct {
client *http.Client
url string
}{
"platform": {
client: NewHTTPClient(time.Second),
url: "https://open.feishu.cn/open-apis/test",
},
"external": {
client: NewHTTPClient(time.Second),
url: "https://external.example/test",
},
"forced external": {
client: NewExternalHTTPClient(time.Second),
url: "https://external.example/test",
},
} {
t.Run(name, func(t *testing.T) {
resp, err := test.client.Get(test.url)
if err == nil {
if resp != nil && resp.Body != nil {
resp.Body.Close()
}
t.Fatalf("policy-routed client succeeded with malformed proxy config")
}
})
}
}

View File

@@ -62,10 +62,7 @@ func httpClient() *http.Client {
if DefaultClient != nil {
return DefaultClient
}
return &http.Client{
Timeout: fetchTimeout,
Transport: transport.Shared(),
}
return transport.NewExternalHTTPClient(fetchTimeout)
}
// updateState is persisted to disk for caching.

View File

@@ -4,6 +4,7 @@
package update
import (
"context"
"encoding/json"
"net/http"
"net/http/httptest"
@@ -12,6 +13,8 @@ import (
"path/filepath"
"testing"
"time"
exttransport "github.com/larksuite/cli/extension/transport"
)
// roundTripFunc adapts a function to http.RoundTripper.
@@ -19,6 +22,30 @@ type roundTripFunc func(*http.Request) (*http.Response, error)
func (f roundTripFunc) RoundTrip(req *http.Request) (*http.Response, error) { return f(req) }
type updateExternalProvider struct {
interceptor exttransport.Interceptor
}
func (p updateExternalProvider) Name() string { return "update-external-test" }
func (p updateExternalProvider) ResolveInterceptor(context.Context) exttransport.Interceptor {
return p.interceptor
}
func (updateExternalProvider) SupportsRequestClass(class exttransport.RequestClass) bool {
return class == exttransport.RequestClassExternal
}
type updateExternalInterceptor struct {
calls int
}
func (i *updateExternalInterceptor) PreRoundTrip(req *http.Request) func(*http.Response, error) {
i.calls++
req.Header.Set("X-External-Route", "1")
return nil
}
// clearSkipEnv unsets all env vars that shouldSkip checks,
// preventing the host environment (e.g. CI=true) from polluting test results.
func clearSkipEnv(t *testing.T) {
@@ -242,6 +269,46 @@ func TestRefreshCache(t *testing.T) {
RefreshCache("1.0.0")
}
func TestHTTPClientUsesExternalRequestClass(t *testing.T) {
t.Setenv("LARKSUITE_CLI_CONFIG_DIR", t.TempDir())
t.Setenv("LARK_CLI_NO_PROXY", "")
previousClient := DefaultClient
DefaultClient = nil
t.Cleanup(func() { DefaultClient = previousClient })
previousProvider := exttransport.GetProvider()
interceptor := &updateExternalInterceptor{}
exttransport.Register(updateExternalProvider{interceptor: interceptor})
t.Cleanup(func() { exttransport.Register(previousProvider) })
previousTransport := http.DefaultTransport
var receivedHeader string
http.DefaultTransport = roundTripFunc(func(req *http.Request) (*http.Response, error) {
receivedHeader = req.Header.Get("X-External-Route")
return &http.Response{
StatusCode: http.StatusNoContent,
Header: make(http.Header),
Body: http.NoBody,
Request: req,
}, nil
})
t.Cleanup(func() { http.DefaultTransport = previousTransport })
req, err := http.NewRequest(http.MethodGet, "https://open.feishu.cn/npm/latest", nil)
if err != nil {
t.Fatal(err)
}
resp, err := httpClient().Do(req)
if err != nil {
t.Fatal(err)
}
resp.Body.Close()
if interceptor.calls != 1 || receivedHeader != "1" {
t.Fatalf("external route = calls %d, header %q; want 1, %q", interceptor.calls, receivedHeader, "1")
}
}
func TestPendingAtomicAccess(t *testing.T) {
// Initially nil
if got := GetPending(); got != nil {

View File

@@ -5,11 +5,15 @@ package validate
import (
"context"
"crypto/tls"
"fmt"
"net"
"net/http"
"net/url"
"strings"
"sync"
"github.com/larksuite/cli/errs"
)
const (
@@ -34,6 +38,9 @@ func isRestrictedDownloadIP(ip net.IP) bool {
return true
}
if v4 := ip.To4(); v4 != nil {
if v4[0] == 0 { // RFC 1122 "this network"
return true
}
if v4[0] == 10 || v4[0] == 127 {
return true
}
@@ -52,6 +59,9 @@ func isRestrictedDownloadIP(ip net.IP) bool {
if v4[0] == 198 && (v4[1] == 18 || v4[1] == 19) { // RFC2544 benchmarking
return true
}
if v4[0] >= 240 {
return true
}
return false
}
if ip.IsPrivate() {
@@ -76,32 +86,42 @@ func ValidateDownloadSourceURL(ctx context.Context, rawURL string) error {
if u.Scheme != "http" && u.Scheme != "https" {
return fmt.Errorf("only http/https URLs are supported")
}
host := strings.TrimSpace(strings.ToLower(u.Hostname()))
_, err = resolveDownloadHost(ctx, u.Hostname(), net.DefaultResolver.LookupIP)
return err
}
type downloadLookupIPFunc func(context.Context, string, string) ([]net.IP, error)
func resolveDownloadHost(ctx context.Context, rawHost string, lookupIP downloadLookupIPFunc) ([]net.IP, error) {
host := strings.TrimSpace(strings.ToLower(rawHost))
if host == "" {
return fmt.Errorf("URL host is required")
return nil, fmt.Errorf("URL host is required")
}
if host == "localhost" || strings.HasSuffix(host, ".localhost") {
return fmt.Errorf("local/internal host is not allowed")
return nil, fmt.Errorf("local/internal host is not allowed")
}
if ip := net.ParseIP(host); ip != nil {
if isRestrictedDownloadIP(ip) {
return fmt.Errorf("local/internal host is not allowed")
return nil, fmt.Errorf("local/internal host is not allowed")
}
return nil
return []net.IP{ip}, nil
}
ips, err := net.DefaultResolver.LookupIP(ctx, "ip", host)
if lookupIP == nil {
lookupIP = net.DefaultResolver.LookupIP
}
ips, err := lookupIP(ctx, "ip", host)
if err != nil {
return fmt.Errorf("failed to resolve host")
return nil, fmt.Errorf("failed to resolve host")
}
if len(ips) == 0 {
return fmt.Errorf("failed to resolve host")
return nil, fmt.Errorf("failed to resolve host")
}
for _, ip := range ips {
if isRestrictedDownloadIP(ip) {
return fmt.Errorf("local/internal host is not allowed")
return nil, fmt.Errorf("local/internal host is not allowed")
}
}
return nil
return ips, nil
}
// NewDownloadHTTPClient clones base client and enforces download-safe redirect
@@ -115,7 +135,10 @@ func NewDownloadHTTPClient(base *http.Client, opts DownloadHTTPClientOptions) *h
}
cloned := *base
cloned.Transport = cloneDownloadTransport(base.Transport)
cloned.Transport = &downloadSchemeTransport{
base: cloneDownloadTransport(base.Transport),
allowHTTP: opts.AllowHTTP,
}
cloned.CheckRedirect = func(req *http.Request, via []*http.Request) error {
if len(via) >= opts.MaxRedirects {
return fmt.Errorf("too many redirects")
@@ -138,18 +161,310 @@ func NewDownloadHTTPClient(base *http.Client, opts DownloadHTTPClientOptions) *h
return &cloned
}
func cloneDownloadTransport(base http.RoundTripper) *http.Transport {
var cloned *http.Transport
if src, ok := base.(*http.Transport); ok && src != nil {
cloned = src.Clone()
} else {
if def, ok := http.DefaultTransport.(*http.Transport); ok && def != nil {
cloned = def.Clone()
} else {
cloned = &http.Transport{}
}
type downloadSchemeTransport struct {
base http.RoundTripper
allowHTTP bool
}
func (t *downloadSchemeTransport) RoundTrip(req *http.Request) (*http.Response, error) {
if req == nil || req.URL == nil {
return nil, errs.NewInternalError(
errs.SubtypeUnknown,
"download transport received a nil request",
)
}
switch {
case strings.EqualFold(req.URL.Scheme, "https"):
case t.allowHTTP && strings.EqualFold(req.URL.Scheme, "http"):
default:
return nil, errs.NewSecurityPolicyError(
errs.SubtypeAccessDenied,
"only https URLs are supported",
)
}
return t.base.RoundTrip(req)
}
type selectedDownloadProxyKey struct{}
type proxyAwareDownloadTransport struct {
selectProxy func(*http.Request) (*url.URL, error)
direct http.RoundTripper
proxied *http.Transport
lookupIP downloadLookupIPFunc
mu sync.Mutex
proxiedByTLSServer map[string]*http.Transport
}
func (t *proxyAwareDownloadTransport) RoundTrip(req *http.Request) (*http.Response, error) {
if req == nil || req.URL == nil {
return nil, fmt.Errorf("download transport received a nil request")
}
proxyURL, err := t.selectProxy(req)
if err != nil {
return nil, err
}
if proxyURL == nil {
return t.direct.RoundTrip(req)
}
targetIPs, err := resolveDownloadHost(req.Context(), req.URL.Hostname(), t.lookupIP)
if err != nil {
return nil, errs.NewSecurityPolicyError(
errs.SubtypeAccessDenied,
"blocked download target: %v",
err,
).WithCause(err)
}
if strings.EqualFold(req.URL.Scheme, "http") && net.ParseIP(req.URL.Hostname()) == nil {
// HTTP proxies cannot pin the target IP separately from the Host header.
return nil, errs.NewSecurityPolicyError(
errs.SubtypeAccessDenied,
"plain HTTP hostname downloads through a proxy are not allowed",
).WithHint("use HTTPS or a literal public IP")
}
selected := *proxyURL
proxied := t.proxied
if strings.EqualFold(req.URL.Scheme, "https") {
proxied = t.proxiedTransportForTLSServer(req.URL.Hostname())
}
var lastErr error
for index, targetIP := range targetIPs {
proxiedReq, pinErr := pinDownloadRequestTargetToIP(req, targetIP)
if pinErr != nil {
return nil, pinErr
}
ctx := context.WithValue(proxiedReq.Context(), selectedDownloadProxyKey{}, &selected)
proxiedReq = proxiedReq.WithContext(ctx)
resp, roundTripErr := proxied.RoundTrip(proxiedReq)
if roundTripErr == nil {
if resp != nil {
// Hide the internal pinned URL from redirect handling.
resp.Request = req
}
return resp, nil
}
lastErr = roundTripErr
if resp != nil && resp.Body != nil {
resp.Body.Close()
}
if req.Context().Err() != nil {
break
}
if index+1 < len(targetIPs) && !canRetryDownloadTarget(req) {
break
}
}
return nil, lastErr
}
func (t *proxyAwareDownloadTransport) CloseIdleConnections() {
if closer, ok := t.direct.(interface{ CloseIdleConnections() }); ok {
closer.CloseIdleConnections()
}
t.proxied.CloseIdleConnections()
t.mu.Lock()
defer t.mu.Unlock()
for _, transport := range t.proxiedByTLSServer {
transport.CloseIdleConnections()
}
}
func (t *proxyAwareDownloadTransport) proxiedTransportForTLSServer(serverName string) *http.Transport {
if configured := t.proxied.TLSClientConfig; configured != nil && configured.ServerName != "" {
serverName = configured.ServerName
}
t.mu.Lock()
defer t.mu.Unlock()
if transport := t.proxiedByTLSServer[serverName]; transport != nil {
return transport
}
transport := t.proxied.Clone()
targetTLSConfig := cloneDownloadTLSConfig(transport.TLSClientConfig)
targetTLSConfig.ServerName = serverName
transport.TLSClientConfig = targetTLSConfig
configureHTTPSProxyTLSDialer(transport, t.proxied)
if t.proxiedByTLSServer == nil {
t.proxiedByTLSServer = make(map[string]*http.Transport)
}
t.proxiedByTLSServer[serverName] = transport
return transport
}
type blockedDownloadTransport struct {
err error
}
func (t *blockedDownloadTransport) RoundTrip(*http.Request) (*http.Response, error) {
return nil, t.err
}
func cloneDownloadTransport(base http.RoundTripper) http.RoundTripper {
if base == nil {
base = http.DefaultTransport
}
if source, ok := base.(interface {
TransformHTTPTransport(func(*http.Transport) (http.RoundTripper, bool)) (http.RoundTripper, bool)
}); ok {
rebuilt, transformed := source.TransformHTTPTransport(newDownloadTransportLeaf)
if transformed && rebuilt != nil {
return rebuilt
}
}
if source, ok := base.(*http.Transport); ok && source != nil {
rebuilt, transformed := newDownloadTransportLeaf(source)
if transformed && rebuilt != nil {
return rebuilt
}
}
return &blockedDownloadTransport{err: errs.NewInternalError(
errs.SubtypeUnknown,
"cannot safely clone download transport %T",
base,
)}
}
func newDownloadTransportLeaf(source *http.Transport) (http.RoundTripper, bool) {
return newDownloadTransportLeafWithResolver(source, net.DefaultResolver.LookupIP)
}
func newDownloadTransportLeafWithResolver(source *http.Transport, lookupIP downloadLookupIPFunc) (http.RoundTripper, bool) {
if source == nil {
return nil, false
}
selectProxy := source.Proxy
direct := cloneDownloadHTTPTransport(source)
direct.Proxy = nil
configureDirectDownloadTransport(direct)
if selectProxy == nil {
return direct, true
}
// The proxied branch validates the requested URL before construction and
// on every redirect. Its TCP peer is the selected proxy, so applying the
// direct-origin IP guard there would incorrectly reject trusted loopback or
// private-network proxies. Freeze the selected proxy in request context so
// a stateful selector cannot switch the second lookup to direct egress.
proxied := cloneDownloadHTTPTransport(source)
proxied.Proxy = func(req *http.Request) (*url.URL, error) {
selected, ok := req.Context().Value(selectedDownloadProxyKey{}).(*url.URL)
if !ok || selected == nil {
return nil, fmt.Errorf("download proxy selection is missing")
}
cloned := *selected
return &cloned, nil
}
return &proxyAwareDownloadTransport{
selectProxy: selectProxy,
direct: direct,
proxied: proxied,
lookupIP: lookupIP,
proxiedByTLSServer: make(map[string]*http.Transport),
}, true
}
func cloneDownloadHTTPTransport(source *http.Transport) *http.Transport {
cloned := source.Clone()
if cloned.TLSNextProto == nil {
if _, ok := source.TLSNextProto["h2"]; ok {
cloned.ForceAttemptHTTP2 = true
}
}
return cloned
}
func pinDownloadRequestTargetToIP(req *http.Request, targetIP net.IP) (*http.Request, error) {
if req == nil || req.URL == nil {
return nil, fmt.Errorf("download request URL is missing")
}
if targetIP == nil || isRestrictedDownloadIP(targetIP) {
return nil, fmt.Errorf("blocked download target: local/internal host is not allowed")
}
originalHost := req.URL.Host
pinnedHost := targetIP.String()
if port := req.URL.Port(); port != "" {
pinnedHost = net.JoinHostPort(pinnedHost, port)
} else if strings.Contains(pinnedHost, ":") {
pinnedHost = "[" + pinnedHost + "]"
}
pinned := req.Clone(req.Context())
pinnedURL := *req.URL
pinnedURL.Host = pinnedHost
pinned.URL = &pinnedURL
pinned.Host = originalHost
return pinned, nil
}
func canRetryDownloadTarget(req *http.Request) bool {
if req == nil || req.Body != nil {
return false
}
return req.Method == http.MethodGet || req.Method == http.MethodHead
}
func cloneDownloadTLSConfig(config *tls.Config) *tls.Config {
if config == nil {
return &tls.Config{MinVersion: tls.VersionTLS12}
}
return config.Clone()
}
func configureHTTPSProxyTLSDialer(transport, source *http.Transport) {
if transport.DialTLSContext != nil || transport.DialTLS != nil {
return
}
proxyTLSConfig := cloneDownloadTLSConfig(source.TLSClientConfig)
transport.DialTLSContext = func(ctx context.Context, network, addr string) (net.Conn, error) {
rawConn, err := dialDownloadProxy(ctx, source, network, addr)
if err != nil {
return nil, err
}
config := proxyTLSConfig.Clone()
serverName, _, splitErr := net.SplitHostPort(addr)
if splitErr != nil {
rawConn.Close()
return nil, fmt.Errorf("invalid HTTPS proxy address: %w", splitErr)
}
config.ServerName = serverName
tlsConn := tls.Client(rawConn, config)
handshakeCtx := ctx
cancel := func() {}
if source.TLSHandshakeTimeout > 0 {
handshakeCtx, cancel = context.WithTimeout(ctx, source.TLSHandshakeTimeout)
}
defer cancel()
if err := tlsConn.HandshakeContext(handshakeCtx); err != nil {
rawConn.Close()
return nil, err
}
return tlsConn, nil
}
}
func dialDownloadProxy(ctx context.Context, source *http.Transport, network, addr string) (net.Conn, error) {
if source.DialContext != nil {
return source.DialContext(ctx, network, addr)
}
if source.Dial != nil {
return source.Dial(network, addr)
}
var dialer net.Dialer
return dialer.DialContext(ctx, network, addr)
}
func configureDirectDownloadTransport(cloned *http.Transport) {
origDial := cloned.DialContext
cloned.DialContext = func(ctx context.Context, network, addr string) (net.Conn, error) {
conn, err := dialConn(ctx, origDial, network, addr)
@@ -158,7 +473,7 @@ func cloneDownloadTransport(base http.RoundTripper) *http.Transport {
}
if err := validateConnRemoteIP(conn); err != nil {
conn.Close()
return nil, err
return nil, downloadTargetPolicyError(err)
}
return conn, nil
}
@@ -172,13 +487,26 @@ func cloneDownloadTransport(base http.RoundTripper) *http.Transport {
}
if err := validateConnRemoteIP(conn); err != nil {
conn.Close()
return nil, downloadTargetPolicyError(err)
}
return conn, nil
}
}
if cloned.DialTLS != nil {
origDialTLS := cloned.DialTLS
cloned.DialTLS = func(network, addr string) (net.Conn, error) {
conn, err := origDialTLS(network, addr)
if err != nil {
return nil, err
}
if err := validateConnRemoteIP(conn); err != nil {
conn.Close()
return nil, downloadTargetPolicyError(err)
}
return conn, nil
}
}
return cloned
}
// DialContextFunc is the signature for DialContext / DialTLSContext.
@@ -194,7 +522,7 @@ func WrapDialContextWithIPCheck(origDial DialContextFunc) DialContextFunc {
}
if err := validateConnRemoteIP(conn); err != nil {
conn.Close()
return nil, err
return nil, downloadTargetPolicyError(err)
}
return conn, nil
}
@@ -208,6 +536,14 @@ func dialConn(ctx context.Context, dialFn func(context.Context, string, string)
return d.DialContext(ctx, network, addr)
}
func downloadTargetPolicyError(err error) error {
return errs.NewSecurityPolicyError(
errs.SubtypeAccessDenied,
"blocked download target: %v",
err,
).WithCause(err)
}
func validateConnRemoteIP(conn net.Conn) error {
if conn == nil {
return fmt.Errorf("nil connection")

View File

@@ -0,0 +1,529 @@
// Copyright (c) 2026 Lark Technologies Pte. Ltd.
// SPDX-License-Identifier: MIT
package validate
import (
"context"
"crypto/tls"
"crypto/x509"
"errors"
"net"
"net/http"
"net/http/httptest"
"net/url"
"strings"
"sync/atomic"
"testing"
"time"
"github.com/larksuite/cli/errs"
)
func TestProxiedHTTPSDownloadPinsValidatedTargetIP(t *testing.T) {
const (
targetHost = "rebind.example"
targetIP = "203.0.113.10"
)
proxyCalled := make(chan struct{}, 1)
proxy := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, req *http.Request) {
proxyCalled <- struct{}{}
if req.Method != http.MethodConnect {
t.Errorf("proxy request method = %q, want CONNECT", req.Method)
}
if got := req.Host; got != targetIP+":443" {
t.Errorf("proxy CONNECT target = %q, want validated IP %q", got, targetIP+":443")
}
w.WriteHeader(http.StatusBadGateway)
}))
t.Cleanup(proxy.Close)
proxyURL, err := url.Parse(proxy.URL)
if err != nil {
t.Fatal(err)
}
lookupIP := func(context.Context, string, string) ([]net.IP, error) {
return []net.IP{net.ParseIP(targetIP)}, nil
}
transport, ok := newDownloadTransportLeafWithResolver(
&http.Transport{Proxy: http.ProxyURL(proxyURL)},
lookupIP,
)
if !ok {
t.Fatal("newDownloadTransportLeafWithResolver() did not rebuild transport")
}
req, err := http.NewRequest(http.MethodGet, "https://"+targetHost+"/file", nil)
if err != nil {
t.Fatal(err)
}
pinned, err := pinDownloadRequestTargetToIP(req, net.ParseIP(targetIP))
if err != nil {
t.Fatal(err)
}
if pinned.Host != targetHost {
t.Fatalf("pinned request Host = %q, want %q", pinned.Host, targetHost)
}
if _, err := transport.RoundTrip(req); err == nil {
t.Fatal("RoundTrip() error = nil, want proxy rejection after CONNECT")
}
select {
case <-proxyCalled:
default:
t.Fatal("proxy was not called")
}
}
func TestRestrictedDownloadIPBlocksReservedIPv4(t *testing.T) {
for _, rawIP := range []string{"0.1.2.3", "240.0.0.1"} {
if !isRestrictedDownloadIP(net.ParseIP(rawIP)) {
t.Fatalf("%s was classified as safe", rawIP)
}
}
if isRestrictedDownloadIP(net.ParseIP("1.1.1.1")) {
t.Fatal("1.1.1.1 was classified as restricted")
}
}
func TestCloneDownloadTLSConfigSetsMinimumVersion(t *testing.T) {
if got := cloneDownloadTLSConfig(nil).MinVersion; got != tls.VersionTLS12 {
t.Fatalf("MinVersion = %d, want TLS 1.2", got)
}
configured := &tls.Config{MinVersion: tls.VersionTLS13}
if got := cloneDownloadTLSConfig(configured).MinVersion; got != tls.VersionTLS13 {
t.Fatalf("cloned MinVersion = %d, want TLS 1.3", got)
}
}
func TestCloneDownloadHTTPTransportPreservesHTTP2Policy(t *testing.T) {
tests := []struct {
name string
source *http.Transport
wantForce bool
wantH2Handler bool
wantProtocolMap bool
}{
{
name: "automatic",
source: &http.Transport{
Proxy: http.ProxyURL(&url.URL{Scheme: "http", Host: "proxy.example:8080"}),
},
wantForce: true,
},
{
name: "custom TLS without opt-in",
source: &http.Transport{
TLSClientConfig: &tls.Config{MinVersion: tls.VersionTLS12},
},
},
{
name: "custom dial without opt-in",
source: &http.Transport{
DialContext: func(context.Context, string, string) (net.Conn, error) {
return nil, errors.New("unused")
},
},
},
{
name: "explicit opt-in",
source: &http.Transport{ForceAttemptHTTP2: true},
wantForce: true,
},
{
name: "explicit h2 handler",
source: &http.Transport{
TLSNextProto: map[string]func(string, *tls.Conn) http.RoundTripper{
"h2": func(string, *tls.Conn) http.RoundTripper { return nil },
},
},
wantH2Handler: true,
wantProtocolMap: true,
},
{
name: "explicit opt-out",
source: &http.Transport{
TLSNextProto: map[string]func(string, *tls.Conn) http.RoundTripper{},
},
wantProtocolMap: true,
},
}
for _, test := range tests {
t.Run(test.name, func(t *testing.T) {
cloned := cloneDownloadHTTPTransport(test.source)
if cloned.ForceAttemptHTTP2 != test.wantForce {
t.Fatalf("ForceAttemptHTTP2 = %v, want %v", cloned.ForceAttemptHTTP2, test.wantForce)
}
_, hasH2Handler := cloned.TLSNextProto["h2"]
if hasH2Handler != test.wantH2Handler {
t.Fatalf("h2 handler = %v, want %v", hasH2Handler, test.wantH2Handler)
}
if hasProtocolMap := cloned.TLSNextProto != nil; hasProtocolMap != test.wantProtocolMap {
t.Fatalf("TLSNextProto is non-nil = %v, want %v", hasProtocolMap, test.wantProtocolMap)
}
})
}
}
func TestCloneDownloadTransportPreservesAutomaticHTTP2(t *testing.T) {
source := &http.Transport{
Proxy: http.ProxyURL(&url.URL{Scheme: "http", Host: "proxy.example:8080"}),
}
rebuilt := cloneDownloadTransport(source)
proxyAware, ok := rebuilt.(*proxyAwareDownloadTransport)
if !ok {
t.Fatalf("transport type = %T, want *proxyAwareDownloadTransport", rebuilt)
}
direct, ok := proxyAware.direct.(*http.Transport)
if !ok {
t.Fatalf("direct transport type = %T, want *http.Transport", proxyAware.direct)
}
for name, transport := range map[string]*http.Transport{
"direct": direct,
"proxied": proxyAware.proxied,
} {
if !transport.ForceAttemptHTTP2 {
t.Fatalf("%s ForceAttemptHTTP2 = false, want true", name)
}
}
}
func TestProxiedDownloadRejectsRestrictedResolvedTarget(t *testing.T) {
var proxyCalled atomic.Bool
proxy := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) {
proxyCalled.Store(true)
w.WriteHeader(http.StatusNoContent)
}))
t.Cleanup(proxy.Close)
proxyURL, err := url.Parse(proxy.URL)
if err != nil {
t.Fatal(err)
}
lookupIP := func(context.Context, string, string) ([]net.IP, error) {
return []net.IP{net.ParseIP("127.0.0.1")}, nil
}
transport, ok := newDownloadTransportLeafWithResolver(
&http.Transport{Proxy: http.ProxyURL(proxyURL)},
lookupIP,
)
if !ok {
t.Fatal("newDownloadTransportLeafWithResolver() did not rebuild transport")
}
req, err := http.NewRequest(http.MethodGet, "http://rebind.example/file", nil)
if err != nil {
t.Fatal(err)
}
_, err = transport.RoundTrip(req)
if err == nil {
t.Fatal("RoundTrip() error = nil, want restricted target rejection")
}
if problem, ok := errs.ProblemOf(err); !ok ||
problem.Category != errs.CategoryPolicy ||
problem.Subtype != errs.SubtypeAccessDenied {
t.Fatalf("RoundTrip() problem = %#v, %v; want policy/access_denied", problem, ok)
}
if proxyCalled.Load() {
t.Fatal("proxy was called for a restricted resolved target")
}
}
func TestProxiedPlainHTTPHostnameRejectsLocalProxy(t *testing.T) {
var proxyCalled atomic.Bool
proxy := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) {
proxyCalled.Store(true)
w.WriteHeader(http.StatusNoContent)
}))
t.Cleanup(proxy.Close)
proxyURL, err := url.Parse(proxy.URL)
if err != nil {
t.Fatal(err)
}
lookupIP := func(ctx context.Context, _, host string) ([]net.IP, error) {
if host == "public.example" {
return []net.IP{net.ParseIP("203.0.113.10")}, nil
}
return net.DefaultResolver.LookupIP(ctx, "ip", host)
}
transport, ok := newDownloadTransportLeafWithResolver(
&http.Transport{Proxy: http.ProxyURL(proxyURL)},
lookupIP,
)
if !ok {
t.Fatal("newDownloadTransportLeafWithResolver() did not rebuild transport")
}
req, err := http.NewRequest(http.MethodGet, "http://public.example/file", nil)
if err != nil {
t.Fatal(err)
}
_, err = transport.RoundTrip(req)
if err == nil {
t.Fatal("RoundTrip() error = nil, want plain HTTP hostname rejection")
}
if problem, ok := errs.ProblemOf(err); !ok ||
problem.Category != errs.CategoryPolicy ||
problem.Subtype != errs.SubtypeAccessDenied {
t.Fatalf("RoundTrip() problem = %#v, %v; want policy/access_denied", problem, ok)
} else if problem.Hint != "use HTTPS or a literal public IP" {
t.Fatalf("RoundTrip() hint = %q, want recovery guidance", problem.Hint)
}
if proxyCalled.Load() {
t.Fatal("proxy was called for a plain HTTP hostname target")
}
}
func TestProxiedHTTPSDownloadTriesEveryValidatedTargetIP(t *testing.T) {
connectTargets := make(chan string, 2)
proxy := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, req *http.Request) {
connectTargets <- req.Host
w.WriteHeader(http.StatusBadGateway)
}))
t.Cleanup(proxy.Close)
proxyURL, err := url.Parse(proxy.URL)
if err != nil {
t.Fatal(err)
}
transport, ok := newDownloadTransportLeafWithResolver(
&http.Transport{Proxy: http.ProxyURL(proxyURL)},
func(context.Context, string, string) ([]net.IP, error) {
return []net.IP{
net.ParseIP("203.0.113.10"),
net.ParseIP("203.0.113.11"),
}, nil
},
)
if !ok {
t.Fatal("newDownloadTransportLeafWithResolver() did not rebuild transport")
}
req, err := http.NewRequest(http.MethodGet, "https://multi.example/file", nil)
if err != nil {
t.Fatal(err)
}
if _, err := transport.RoundTrip(req); err == nil {
t.Fatal("RoundTrip() error = nil, want proxy rejection")
}
for _, want := range []string{"203.0.113.10:443", "203.0.113.11:443"} {
select {
case got := <-connectTargets:
if got != want {
t.Fatalf("proxy CONNECT target = %q, want %q", got, want)
}
default:
t.Fatalf("proxy did not receive CONNECT target %q", want)
}
}
}
func TestCanRetryDownloadTargetOnlyAllowsBodylessReads(t *testing.T) {
for _, test := range []struct {
method string
body string
want bool
}{
{method: http.MethodGet, want: true},
{method: http.MethodHead, want: true},
{method: http.MethodPost},
{method: http.MethodGet, body: "body"},
} {
req, err := http.NewRequest(test.method, "https://download.example/file", strings.NewReader(test.body))
if err != nil {
t.Fatal(err)
}
if test.body == "" {
req.Body = nil
}
if got := canRetryDownloadTarget(req); got != test.want {
t.Fatalf("canRetryDownloadTarget(%s, body=%q) = %v, want %v", test.method, test.body, got, test.want)
}
}
}
func TestProxiedHTTPSTargetPreservesOriginalTLSServerName(t *testing.T) {
transport, ok := newDownloadTransportLeafWithResolver(
&http.Transport{Proxy: http.ProxyURL(&url.URL{Scheme: "http", Host: "proxy.example:8080"})},
func(context.Context, string, string) ([]net.IP, error) {
return []net.IP{net.ParseIP("203.0.113.10")}, nil
},
)
if !ok {
t.Fatal("newDownloadTransportLeafWithResolver() did not rebuild transport")
}
proxyAware, ok := transport.(*proxyAwareDownloadTransport)
if !ok {
t.Fatalf("transport type = %T, want *proxyAwareDownloadTransport", transport)
}
pinned := proxyAware.proxiedTransportForTLSServer("download.example")
if pinned.TLSClientConfig == nil {
t.Fatal("TLSClientConfig = nil")
}
if pinned.TLSClientConfig.ServerName != "download.example" {
t.Fatalf("TLS ServerName = %q, want download.example", pinned.TLSClientConfig.ServerName)
}
if proxyAware.proxied.TLSClientConfig != nil && proxyAware.proxied.TLSClientConfig.ServerName != "" {
t.Fatalf("base proxy TLS ServerName = %q, want unchanged", proxyAware.proxied.TLSClientConfig.ServerName)
}
}
func TestHTTPSProxyTLSDialerUsesLegacyDial(t *testing.T) {
wantErr := errors.New("legacy dial used")
source := &http.Transport{
Dial: func(string, string) (net.Conn, error) {
return nil, wantErr
},
}
target := source.Clone()
configureHTTPSProxyTLSDialer(target, source)
if target.DialTLSContext == nil {
t.Fatal("DialTLSContext = nil")
}
if _, err := target.DialTLSContext(context.Background(), "tcp", "proxy.example:443"); !errors.Is(err, wantErr) {
t.Fatalf("DialTLSContext() error = %v, want %v", err, wantErr)
}
}
func TestDirectDownloadLegacyDialTLSClosesRestrictedConnection(t *testing.T) {
clientConn, serverConn := net.Pipe()
t.Cleanup(func() { serverConn.Close() })
conn := &trackedDownloadConn{
Conn: clientConn,
remoteAddr: &net.TCPAddr{IP: net.ParseIP("127.0.0.1"), Port: 443},
}
rebuilt, ok := newDownloadTransportLeaf(&http.Transport{
DialTLS: func(string, string) (net.Conn, error) {
return conn, nil
},
})
if !ok {
t.Fatal("newDownloadTransportLeaf() did not rebuild transport")
}
transport, ok := rebuilt.(*http.Transport)
if !ok {
t.Fatalf("rebuilt transport = %T, want *http.Transport", rebuilt)
}
_, err := transport.DialTLS("tcp", "public.example:443")
if err == nil || !strings.Contains(err.Error(), "local/internal host is not allowed") {
t.Fatalf("DialTLS() error = %v, want restricted target rejection", err)
}
if problem, ok := errs.ProblemOf(err); !ok ||
problem.Category != errs.CategoryPolicy ||
problem.Subtype != errs.SubtypeAccessDenied {
t.Fatalf("DialTLS() problem = %#v, %v; want policy/access_denied", problem, ok)
}
if !conn.closed {
t.Fatal("restricted connection was not closed")
}
}
func TestDirectDownloadLegacyDialTLSPreservesDialError(t *testing.T) {
wantErr := errors.New("dial failed")
rebuilt, ok := newDownloadTransportLeaf(&http.Transport{
DialTLS: func(string, string) (net.Conn, error) {
return nil, wantErr
},
})
if !ok {
t.Fatal("newDownloadTransportLeaf() did not rebuild transport")
}
transport, ok := rebuilt.(*http.Transport)
if !ok {
t.Fatalf("rebuilt transport = %T, want *http.Transport", rebuilt)
}
if _, err := transport.DialTLS("tcp", "public.example:443"); !errors.Is(err, wantErr) {
t.Fatalf("DialTLS() error = %v, want %v", err, wantErr)
}
}
func TestHTTPSProxyTLSDialerRetainsHandshakeTimeout(t *testing.T) {
clientConn, serverConn := net.Pipe()
t.Cleanup(func() {
clientConn.Close()
serverConn.Close()
})
source := &http.Transport{
DialContext: func(context.Context, string, string) (net.Conn, error) {
return clientConn, nil
},
TLSHandshakeTimeout: 50 * time.Millisecond,
}
target := source.Clone()
configureHTTPSProxyTLSDialer(target, source)
started := time.Now()
if _, err := target.DialTLSContext(context.Background(), "tcp", "proxy.example:443"); err == nil {
t.Fatal("DialTLSContext() error = nil, want TLS handshake timeout")
}
if elapsed := time.Since(started); elapsed > time.Second {
t.Fatalf("TLS handshake timeout took %s, want under 1s", elapsed)
}
}
func TestHTTPSProxyTLSDialerUsesProxyServerName(t *testing.T) {
server := httptest.NewTLSServer(http.HandlerFunc(func(http.ResponseWriter, *http.Request) {}))
t.Cleanup(server.Close)
clientConn, serverConn := net.Pipe()
t.Cleanup(func() {
clientConn.Close()
serverConn.Close()
})
proxySNI := make(chan string, 1)
serverTLSConfig := server.TLS.Clone()
serverTLSConfig.GetConfigForClient = func(info *tls.ClientHelloInfo) (*tls.Config, error) {
proxySNI <- info.ServerName
return nil, nil
}
serverErr := make(chan error, 1)
go func() {
serverErr <- tls.Server(serverConn, serverTLSConfig).Handshake()
}()
roots := x509.NewCertPool()
roots.AddCert(server.Certificate())
source := &http.Transport{
DialContext: func(context.Context, string, string) (net.Conn, error) {
return clientConn, nil
},
TLSClientConfig: &tls.Config{
RootCAs: roots,
ServerName: "target.example.com",
},
}
target := source.Clone()
configureHTTPSProxyTLSDialer(target, source)
conn, err := target.DialTLSContext(context.Background(), "tcp", "example.com:443")
if err != nil {
t.Fatal(err)
}
conn.Close()
if err := <-serverErr; err != nil {
t.Fatal(err)
}
if got := <-proxySNI; got != "example.com" {
t.Fatalf("proxy TLS ServerName = %q, want example.com", got)
}
}
type trackedDownloadConn struct {
net.Conn
remoteAddr net.Addr
closed bool
}
func (c *trackedDownloadConn) RemoteAddr() net.Addr {
return c.remoteAddr
}
func (c *trackedDownloadConn) Close() error {
c.closed = true
return c.Conn.Close()
}

View File

@@ -0,0 +1,222 @@
// Copyright (c) 2026 Lark Technologies Pte. Ltd.
// SPDX-License-Identifier: MIT
package validate_test
import (
"context"
"crypto/tls"
"errors"
"net"
"net/http"
"net/http/httptest"
"net/url"
"strings"
"testing"
"github.com/larksuite/cli/errs"
exttransport "github.com/larksuite/cli/extension/transport"
internaltransport "github.com/larksuite/cli/internal/transport"
"github.com/larksuite/cli/internal/validate"
)
type opaqueRoundTripper struct {
called bool
}
func (t *opaqueRoundTripper) RoundTrip(req *http.Request) (*http.Response, error) {
t.called = true
return &http.Response{StatusCode: http.StatusNoContent, Body: http.NoBody, Request: req}, nil
}
type downloadTestProvider struct {
interceptor exttransport.Interceptor
}
func (p downloadTestProvider) Name() string { return "download-test" }
func (p downloadTestProvider) ResolveInterceptor(context.Context) exttransport.Interceptor {
return p.interceptor
}
type downloadHeaderInterceptor struct{}
func (downloadHeaderInterceptor) PreRoundTrip(req *http.Request) func(*http.Response, error) {
req.Header.Set("X-Use-Proxy", "1")
return nil
}
func TestNewDownloadHTTPClientPreservesPolicyRouterBaseTransport(t *testing.T) {
wantErr := errors.New("proxy policy blocked request")
base := &http.Transport{
Proxy: func(*http.Request) (*url.URL, error) {
return nil, wantErr
},
}
router := internaltransport.NewHTTPPolicyRouter(base, base)
client := internaltransport.ClientForRequestClass(
&http.Client{Transport: router},
exttransport.RequestClassExternal,
)
download := validate.NewDownloadHTTPClient(client, validate.DownloadHTTPClientOptions{AllowHTTP: true})
req, err := http.NewRequest(http.MethodGet, "https://external.example/file", nil)
if err != nil {
t.Fatal(err)
}
resp, err := download.Transport.RoundTrip(req)
if resp != nil && resp.Body != nil {
resp.Body.Close()
}
if !errors.Is(err, wantErr) {
t.Fatalf("RoundTrip() error = %v, want preserved proxy error %v", err, wantErr)
}
}
func TestNewDownloadHTTPClientRejectsInitialHTTPBeforeTransport(t *testing.T) {
base := &opaqueRoundTripper{}
download := validate.NewDownloadHTTPClient(
&http.Client{Transport: base},
validate.DownloadHTTPClientOptions{},
)
req, err := http.NewRequest(http.MethodGet, "http://203.0.113.10/file", nil)
if err != nil {
t.Fatal(err)
}
_, err = download.Transport.RoundTrip(req)
if err == nil {
t.Fatal("RoundTrip() error = nil, want initial HTTP rejection")
}
problem, ok := errs.ProblemOf(err)
if !ok || problem.Category != errs.CategoryPolicy || problem.Subtype != errs.SubtypeAccessDenied {
t.Fatalf("RoundTrip() problem = %#v, %v; want policy/access_denied", problem, ok)
}
if base.called {
t.Fatal("base transport was called for a disallowed initial HTTP request")
}
}
func TestNewDownloadHTTPClientAllowsSelectedLoopbackProxy(t *testing.T) {
proxy := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, req *http.Request) {
if req.URL.Host != "203.0.113.10" {
t.Errorf("proxy request target = %q, want 203.0.113.10", req.URL.Host)
}
w.WriteHeader(http.StatusNoContent)
}))
t.Cleanup(proxy.Close)
proxyURL, err := url.Parse(proxy.URL)
if err != nil {
t.Fatal(err)
}
base := &http.Transport{Proxy: http.ProxyURL(proxyURL)}
router := internaltransport.NewHTTPPolicyRouter(base, base)
client := internaltransport.ClientForRequestClass(
&http.Client{Transport: router},
exttransport.RequestClassExternal,
)
download := validate.NewDownloadHTTPClient(client, validate.DownloadHTTPClientOptions{AllowHTTP: true})
req, err := http.NewRequest(http.MethodGet, "http://203.0.113.10/file", nil)
if err != nil {
t.Fatal(err)
}
resp, err := download.Do(req)
if err != nil {
t.Fatalf("download through selected loopback proxy: %v", err)
}
resp.Body.Close()
}
func TestNewDownloadHTTPClientSelectsProxyAfterOuterDecorators(t *testing.T) {
proxy := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, req *http.Request) {
if got := req.Header.Get("X-Use-Proxy"); got != "1" {
t.Errorf("proxy received X-Use-Proxy = %q, want 1", got)
}
w.WriteHeader(http.StatusNoContent)
}))
t.Cleanup(proxy.Close)
proxyURL, err := url.Parse(proxy.URL)
if err != nil {
t.Fatal(err)
}
wantErr := errors.New("proxy selector ran before decorators")
base := &http.Transport{Proxy: func(req *http.Request) (*url.URL, error) {
if req.Header.Get("X-Use-Proxy") != "1" {
return nil, wantErr
}
return proxyURL, nil
}}
previousProvider := exttransport.GetProvider()
exttransport.Register(downloadTestProvider{interceptor: downloadHeaderInterceptor{}})
t.Cleanup(func() { exttransport.Register(previousProvider) })
router := internaltransport.NewHTTPPolicyRouter(base, base)
client := internaltransport.ClientForRequestClass(
&http.Client{Transport: router},
exttransport.RequestClassExternal,
)
download := validate.NewDownloadHTTPClient(client, validate.DownloadHTTPClientOptions{AllowHTTP: true})
req, err := http.NewRequest(http.MethodGet, "http://203.0.113.10/file", nil)
if err != nil {
t.Fatal(err)
}
resp, err := download.Do(req)
if err != nil {
t.Fatalf("download through decorator-selected proxy: %v", err)
}
resp.Body.Close()
}
func TestNewDownloadHTTPClientGuardsLegacyDialTLS(t *testing.T) {
server := httptest.NewTLSServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) {
w.WriteHeader(http.StatusNoContent)
}))
t.Cleanup(server.Close)
base := &http.Transport{DialTLS: func(_, _ string) (net.Conn, error) {
return tls.Dial("tcp", server.Listener.Addr().String(), &tls.Config{InsecureSkipVerify: true}) //nolint:gosec // local TLS server verifies the connection guard.
}}
download := validate.NewDownloadHTTPClient(&http.Client{Transport: base}, validate.DownloadHTTPClientOptions{AllowHTTP: true})
req, err := http.NewRequest(http.MethodGet, "https://public.example/file", nil)
if err != nil {
t.Fatal(err)
}
_, err = download.Transport.RoundTrip(req)
if err == nil || !strings.Contains(err.Error(), "local/internal host is not allowed") {
t.Fatalf("RoundTrip() error = %v, want legacy DialTLS IP guard", err)
}
problem, ok := errs.ProblemOf(err)
if !ok || problem.Category != errs.CategoryPolicy || problem.Subtype != errs.SubtypeAccessDenied {
t.Fatalf("RoundTrip() problem = %#v, %v; want policy/access_denied", problem, ok)
}
var policyErr *errs.SecurityPolicyError
if !errors.As(err, &policyErr) || policyErr.Cause == nil {
t.Fatalf("RoundTrip() error = %T, want policy error with cause", err)
}
}
func TestNewDownloadHTTPClientFailsClosedForOpaqueTransport(t *testing.T) {
opaque := &opaqueRoundTripper{}
client := internaltransport.ClientForRequestClass(
&http.Client{Transport: opaque},
exttransport.RequestClassExternal,
)
download := validate.NewDownloadHTTPClient(client, validate.DownloadHTTPClientOptions{AllowHTTP: true})
req, err := http.NewRequest(http.MethodGet, "https://public.example/file", nil)
if err != nil {
t.Fatal(err)
}
_, err = download.Transport.RoundTrip(req)
if err == nil || !strings.Contains(err.Error(), "cannot safely clone download transport") {
t.Fatalf("RoundTrip() error = %v, want fail-closed clone error", err)
}
problem, ok := errs.ProblemOf(err)
if !ok || problem.Category != errs.CategoryInternal || problem.Subtype != errs.SubtypeUnknown {
t.Fatalf("RoundTrip() problem = %#v, %v; want internal/unknown", problem, ok)
}
if opaque.called {
t.Fatal("opaque transport was called after safe cloning failed")
}
}

4
package-lock.json generated
View File

@@ -1,12 +1,12 @@
{
"name": "@larksuite/cli",
"version": "1.0.80",
"version": "1.0.81",
"lockfileVersion": 3,
"requires": true,
"packages": {
"": {
"name": "@larksuite/cli",
"version": "1.0.80",
"version": "1.0.81",
"cpu": [
"x64",
"arm64",

View File

@@ -1,6 +1,6 @@
{
"name": "@larksuite/cli",
"version": "1.0.80",
"version": "1.0.81",
"description": "The official CLI for Lark/Feishu open platform",
"bin": {
"lark-cli": "scripts/run.js"

View File

@@ -14,6 +14,7 @@ import (
"time"
"github.com/larksuite/cli/errs"
"github.com/larksuite/cli/internal/transport"
"github.com/larksuite/cli/internal/validate"
"github.com/larksuite/cli/shortcuts/common"
)
@@ -74,11 +75,9 @@ func normalizeTimestamp(raw string) (string, error) {
return "", errs.NewValidationError(errs.SubtypeInvalidArgument, "invalid timestamp %q (want relative 7d/2h/30s, date 2026-04-15, datetime 2026-04-15T10:00:00, or ISO 8601 with TZ)", s)
}
// newFileTransferClient 直传 / 直下对象存储 presigned URL 用(绕开 Lark 网关,无需 auth、无超时以容纳大文件
//
//nolint:forbidigo // presigned object-storage transfer bypasses the Lark gateway — raw http.Client is required (no Lark auth, no gateway routing); not a Lark API call, so RuntimeContext.DoAPI does not apply.
//nolint:forbidigo // Presigned transfers use the external HTTP policy.
func newFileTransferClient() *http.Client {
return &http.Client{Transport: http.DefaultTransport}
return transport.NewExternalHTTPClient(0)
}
// URL helpers for the file (storage) CLI commands.

View File

@@ -0,0 +1,79 @@
// Copyright (c) 2026 Lark Technologies Pte. Ltd.
// SPDX-License-Identifier: MIT
package apps
import (
"context"
"net/http"
"testing"
exttransport "github.com/larksuite/cli/extension/transport"
)
type appsExternalProvider struct {
interceptor exttransport.Interceptor
}
func (p appsExternalProvider) Name() string { return "apps-external-test" }
func (p appsExternalProvider) ResolveInterceptor(context.Context) exttransport.Interceptor {
return p.interceptor
}
func (appsExternalProvider) SupportsRequestClass(class exttransport.RequestClass) bool {
return class == exttransport.RequestClassExternal
}
type appsExternalInterceptor struct {
calls int
}
func (i *appsExternalInterceptor) PreRoundTrip(req *http.Request) func(*http.Response, error) {
i.calls++
req.Header.Set("X-External-Route", "1")
return nil
}
type appsRoundTripFunc func(*http.Request) (*http.Response, error)
func (f appsRoundTripFunc) RoundTrip(req *http.Request) (*http.Response, error) {
return f(req)
}
func TestFileTransferClientUsesExternalRequestClass(t *testing.T) {
t.Setenv("LARKSUITE_CLI_CONFIG_DIR", t.TempDir())
t.Setenv("LARK_CLI_NO_PROXY", "")
previousProvider := exttransport.GetProvider()
interceptor := &appsExternalInterceptor{}
exttransport.Register(appsExternalProvider{interceptor: interceptor})
t.Cleanup(func() { exttransport.Register(previousProvider) })
previousTransport := http.DefaultTransport
var receivedHeader string
http.DefaultTransport = appsRoundTripFunc(func(req *http.Request) (*http.Response, error) {
receivedHeader = req.Header.Get("X-External-Route")
return &http.Response{
StatusCode: http.StatusNoContent,
Header: make(http.Header),
Body: http.NoBody,
Request: req,
}, nil
})
t.Cleanup(func() { http.DefaultTransport = previousTransport })
req, err := http.NewRequest(http.MethodGet, "https://open.feishu.cn/presigned/file", nil)
if err != nil {
t.Fatal(err)
}
resp, err := newFileTransferClient().Do(req)
if err != nil {
t.Fatal(err)
}
resp.Body.Close()
if interceptor.calls != 1 || receivedHeader != "1" {
t.Fatalf("external route = calls %d, header %q; want 1, %q", interceptor.calls, receivedHeader, "1")
}
}

View File

@@ -40,12 +40,6 @@ func (c docCoverHTTPStatusCause) Error() string {
return http.StatusText(int(c))
}
type docCoverURLGuardError string
func (e docCoverURLGuardError) Error() string {
return string(e)
}
var docCoverAllowedContentTypes = map[string]string{
"image/gif": ".gif",
"image/jpeg": ".jpg",
@@ -542,7 +536,7 @@ func downloadDocCoverURL(ctx context.Context, runtime *common.RuntimeContext, ra
return nil, "", err
}
baseClient, err := runtime.Factory.HttpClient()
baseClient, err := runtime.Factory.ExternalHTTPClient()
if err != nil {
return nil, "", errs.NewInternalError(errs.SubtypeSDKError, "http client: %v", err).WithCause(err)
}
@@ -673,6 +667,9 @@ func isUnsafeDocCoverIP(ip net.IP) bool {
return true
}
if v4 := ip.To4(); v4 != nil {
if v4[0] == 0 {
return true
}
if v4[0] == 10 || v4[0] == 127 {
return true
}
@@ -701,13 +698,15 @@ func isUnsafeDocCoverIP(ip net.IP) bool {
func newDocCoverHTTPClient(base *http.Client) *http.Client { //nolint:forbidigo // guarded external --url downloader cannot use Lark API runtime helpers.
if base == nil {
base = &http.Client{} //nolint:forbidigo // fallback only; caller normally supplies Factory.HttpClient.
base = &http.Client{} //nolint:forbidigo // fallback only; caller normally supplies Factory.ExternalHTTPClient.
}
cloned := *base
if cloned.Timeout == 0 { //nolint:forbidigo // external download timeout guard on cloned client.
cloned.Timeout = 30 * time.Second //nolint:forbidigo // external download timeout guard on cloned client.
}
cloned.Transport = cloneDocCoverTransport(base.Transport) //nolint:forbidigo // external download transport adds proxy/IP guards.
cloned.Transport = validate.NewDownloadHTTPClient(base, validate.DownloadHTTPClientOptions{ //nolint:forbidigo // guarded external download
MaxRedirects: 3,
}).Transport
cloned.CheckRedirect = func(req *http.Request, via []*http.Request) error { //nolint:forbidigo // redirects must be validated for external --url downloads.
if len(via) >= 3 {
return errs.NewValidationError(errs.SubtypeInvalidArgument, "cover URL redirects too many times").WithParam("--url")
@@ -723,73 +722,3 @@ func newDocCoverHTTPClient(base *http.Client) *http.Client { //nolint:forbidigo
}
return &cloned
}
func cloneDocCoverTransport(base http.RoundTripper) *http.Transport { //nolint:forbidigo // external --url downloader wraps caller transport with IP/proxy guards.
var cloned *http.Transport
if src, ok := base.(*http.Transport); ok && src != nil {
cloned = src.Clone()
} else if def, ok := http.DefaultTransport.(*http.Transport); ok && def != nil { //nolint:forbidigo // fallback for guarded external downloader only.
cloned = def.Clone()
} else {
cloned = &http.Transport{}
}
cloned.Proxy = nil
origDial := cloned.DialContext
cloned.DialContext = func(ctx context.Context, network, addr string) (net.Conn, error) {
conn, err := dialDocCoverConn(ctx, origDial, network, addr)
if err != nil {
return nil, err
}
if err := validateDocCoverConnRemoteIP(conn); err != nil {
conn.Close()
return nil, err
}
return conn, nil
}
if cloned.DialTLSContext != nil {
origDialTLS := cloned.DialTLSContext
cloned.DialTLSContext = func(ctx context.Context, network, addr string) (net.Conn, error) {
conn, err := dialDocCoverConn(ctx, origDialTLS, network, addr)
if err != nil {
return nil, err
}
if err := validateDocCoverConnRemoteIP(conn); err != nil {
conn.Close()
return nil, err
}
return conn, nil
}
}
return cloned
}
func dialDocCoverConn(ctx context.Context, dialFn func(context.Context, string, string) (net.Conn, error), network, addr string) (net.Conn, error) {
if dialFn != nil {
return dialFn(ctx, network, addr)
}
var dialer net.Dialer
return dialer.DialContext(ctx, network, addr)
}
func validateDocCoverConnRemoteIP(conn net.Conn) error {
if conn == nil {
return docCoverURLGuardError("nil connection")
}
addr := conn.RemoteAddr()
if addr == nil {
return docCoverURLGuardError("missing remote address")
}
host, _, err := net.SplitHostPort(addr.String())
if err != nil {
host = addr.String()
}
ip := net.ParseIP(strings.Trim(host, "[]"))
if ip == nil {
return docCoverURLGuardError("invalid remote IP")
}
if isUnsafeDocCoverIP(ip) {
return docCoverURLGuardError("local/internal host is not allowed")
}
return nil
}

View File

@@ -386,6 +386,7 @@ func TestValidateDocCoverURLHost(t *testing.T) {
func TestDocCoverIPSafetyBlocksSpecialRanges(t *testing.T) {
for _, rawIP := range []string{
"0.1.2.3",
"10.0.0.1",
"127.0.0.1",
"169.254.1.1",
@@ -406,17 +407,26 @@ func TestDocCoverIPSafetyBlocksSpecialRanges(t *testing.T) {
}
}
func TestDocCoverHTTPClientDoesNotUseProxy(t *testing.T) {
baseTransport := &http.Transport{Proxy: http.ProxyFromEnvironment}
func TestDocCoverHTTPClientPreservesProxyPolicy(t *testing.T) {
proxyErr := errors.New("proxy selected")
directErr := errors.New("direct dialed")
baseTransport := &http.Transport{
Proxy: func(*http.Request) (*url.URL, error) {
return nil, proxyErr
},
DialContext: func(context.Context, string, string) (net.Conn, error) {
return nil, directErr
},
}
baseClient := &http.Client{Transport: baseTransport}
client := newDocCoverHTTPClient(baseClient)
transport, ok := client.Transport.(*http.Transport)
if !ok {
t.Fatalf("client transport = %T, want *http.Transport", client.Transport)
req, err := http.NewRequest(http.MethodGet, "https://203.0.113.10/cover.png", nil)
if err != nil {
t.Fatal(err)
}
if transport.Proxy != nil {
t.Fatal("cover URL downloader must not inherit proxy settings")
if _, err := client.Transport.RoundTrip(req); !errors.Is(err, proxyErr) {
t.Fatalf("RoundTrip() error = %v, want proxy policy error %v", err, proxyErr)
}
if baseTransport.Proxy == nil {
t.Fatal("base transport proxy was mutated")
@@ -446,21 +456,6 @@ func TestDocCoverHTTPClientRedirectValidation(t *testing.T) {
}
}
func TestDocCoverConnRemoteIPValidation(t *testing.T) {
if err := validateDocCoverConnRemoteIP(nil); err == nil {
t.Fatal("expected nil connection error")
}
if err := validateDocCoverConnRemoteIP(docCoverRemoteAddrConn{}); err == nil {
t.Fatal("expected missing remote address error")
}
if err := validateDocCoverConnRemoteIP(docCoverRemoteAddrConn{addr: testAddr("not-ip")}); err == nil {
t.Fatal("expected invalid remote IP error")
}
if err := validateDocCoverConnRemoteIP(docCoverRemoteAddrConn{addr: &net.TCPAddr{IP: net.ParseIP("127.0.0.1"), Port: 443}}); err == nil {
t.Fatal("expected local remote IP error")
}
}
func TestDocCoverURLFileName(t *testing.T) {
cases := []struct {
raw string
@@ -662,16 +657,6 @@ func (c docCoverRemoteAddrConn) RemoteAddr() net.Addr {
return c.addr
}
type testAddr string
func (a testAddr) Network() string {
return "test"
}
func (a testAddr) String() string {
return string(a)
}
type repeatByteReader byte
func (r repeatByteReader) Read(p []byte) (int, error) {
@@ -710,3 +695,61 @@ func decodeDocResourceOutput(t *testing.T, stdout *bytes.Buffer) docResourceOutp
}
return out
}
type opaqueDocCoverTransport struct {
called bool
}
func (t *opaqueDocCoverTransport) RoundTrip(req *http.Request) (*http.Response, error) {
t.called = true
return &http.Response{StatusCode: http.StatusNoContent, Body: http.NoBody, Request: req}, nil
}
func TestNewDocCoverHTTPClientFailsClosedForOpaqueTransport(t *testing.T) {
opaque := &opaqueDocCoverTransport{}
client := newDocCoverHTTPClient(&http.Client{Transport: opaque})
req, err := http.NewRequest(http.MethodGet, "https://public.example/cover.png", nil)
if err != nil {
t.Fatal(err)
}
_, err = client.Transport.RoundTrip(req)
if err == nil || !strings.Contains(err.Error(), "cannot safely clone download transport") {
t.Fatalf("RoundTrip() error = %v, want fail-closed clone error", err)
}
problem, ok := errs.ProblemOf(err)
if !ok || problem.Category != errs.CategoryInternal || problem.Subtype != errs.SubtypeUnknown {
t.Fatalf("RoundTrip() problem = %#v, %v; want internal/unknown", problem, ok)
}
if opaque.called {
t.Fatal("opaque transport was called after safe cloning failed")
}
}
func TestNewDocCoverHTTPClientGuardsLegacyDialTLS(t *testing.T) {
server := httptest.NewTLSServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) {
w.WriteHeader(http.StatusNoContent)
}))
t.Cleanup(server.Close)
base := &http.Transport{DialTLS: func(_, _ string) (net.Conn, error) {
return tls.Dial("tcp", server.Listener.Addr().String(), &tls.Config{InsecureSkipVerify: true}) //nolint:gosec // local TLS server verifies the connection guard.
}}
client := newDocCoverHTTPClient(&http.Client{Transport: base})
req, err := http.NewRequest(http.MethodGet, "https://public.example/cover.png", nil)
if err != nil {
t.Fatal(err)
}
_, err = client.Transport.RoundTrip(req)
if err == nil || !strings.Contains(err.Error(), "local/internal host is not allowed") {
t.Fatalf("RoundTrip() error = %v, want legacy DialTLS IP guard", err)
}
problem, ok := errs.ProblemOf(err)
if !ok || problem.Category != errs.CategoryPolicy || problem.Subtype != errs.SubtypeAccessDenied {
t.Fatalf("RoundTrip() problem = %#v, %v; want policy/access_denied", problem, ok)
}
var policyErr *errs.SecurityPolicyError
if !errors.As(err, &policyErr) || policyErr.Cause == nil {
t.Fatalf("RoundTrip() error = %T, want policy error with cause", err)
}
}

View File

@@ -199,7 +199,7 @@ func startURLDownload(ctx context.Context, runtime *common.RuntimeContext, rawUR
WithCause(err)
}
httpClient, err := runtime.Factory.HttpClient()
httpClient, err := runtime.Factory.ExternalHTTPClient()
if err != nil {
return nil, "", errs.NewInternalError(errs.SubtypeSDKError, "http client: %v", err).WithCause(err)
}

View File

@@ -28,6 +28,7 @@ import (
"github.com/larksuite/cli/internal/cmdutil"
"github.com/larksuite/cli/internal/core"
"github.com/larksuite/cli/internal/credential"
internaltransport "github.com/larksuite/cli/internal/transport"
"github.com/larksuite/cli/shortcuts/common"
)
@@ -43,6 +44,25 @@ func (f shortcutRoundTripFunc) RoundTrip(req *http.Request) (*http.Response, err
return f(req)
}
type shortcutPolicyDecorator struct {
base http.RoundTripper
fn shortcutRoundTripFunc
}
func (t *shortcutPolicyDecorator) BaseRoundTripper() http.RoundTripper {
return t.base
}
func (t *shortcutPolicyDecorator) WithBaseRoundTripper(base http.RoundTripper) http.RoundTripper {
cloned := *t
cloned.base = base
return &cloned
}
func (t *shortcutPolicyDecorator) RoundTrip(req *http.Request) (*http.Response, error) {
return t.fn(req)
}
func shortcutJSONResponse(status int, body interface{}) *http.Response {
b, _ := json.Marshal(body)
return &http.Response{
@@ -901,6 +921,50 @@ func TestStartURLDownloadBlockedURLCarriesParam(t *testing.T) {
}
}
func TestStartURLDownloadUsesExternalRequestClass(t *testing.T) {
platform := &shortcutPolicyDecorator{
base: http.DefaultTransport,
fn: func(req *http.Request) (*http.Response, error) {
return shortcutRawResponse(http.StatusBadGateway, nil, nil), nil
},
}
external := &shortcutPolicyDecorator{
base: http.DefaultTransport,
fn: func(req *http.Request) (*http.Response, error) {
resp := shortcutRawResponse(http.StatusOK, []byte("image"), nil)
resp.Request = req
return resp, nil
},
}
runtime := &common.RuntimeContext{
Factory: &cmdutil.Factory{
HttpClient: func() (*http.Client, error) {
return &http.Client{
Transport: internaltransport.NewHTTPPolicyRouter(platform, external),
}, nil
},
},
}
resp, _, err := startURLDownload(
context.Background(),
runtime,
"https://open.feishu.cn/presigned/image.png",
"--image",
)
if err != nil {
t.Fatalf("startURLDownload() error = %v, want external route", err)
}
defer resp.Body.Close()
body, err := io.ReadAll(resp.Body)
if err != nil {
t.Fatal(err)
}
if got := string(body); got != "image" {
t.Fatalf("download body = %q, want external payload", got)
}
}
// TestResolveLocalMediaImage verifies that resolveLocalMedia can upload an image
// via uploadImageToIM without double path validation.
func TestResolveLocalMediaImage(t *testing.T) {

View File

@@ -1790,7 +1790,7 @@ func downloadAttachmentContent(runtime *common.RuntimeContext, downloadURL strin
return nil, mailInvalidResponseError("attachment download URL has no host")
}
httpClient, err := runtime.Factory.HttpClient()
httpClient, err := runtime.Factory.ExternalHTTPClient()
if err != nil {
return nil, errs.NewInternalError(errs.SubtypeSDKError, "failed to get HTTP client: %v", err).WithCause(err)
}

View File

@@ -9,6 +9,7 @@ import (
"encoding/base64"
"encoding/json"
"fmt"
"io"
"net/http"
"net/http/httptest"
"os"
@@ -20,6 +21,7 @@ import (
"github.com/spf13/cobra"
"github.com/larksuite/cli/internal/cmdutil"
internaltransport "github.com/larksuite/cli/internal/transport"
"github.com/larksuite/cli/internal/vfs/localfileio"
"github.com/larksuite/cli/shortcuts/common"
"github.com/larksuite/cli/shortcuts/mail/emlbuilder"
@@ -396,6 +398,36 @@ func TestDownloadAttachmentContent_NoAuthorizationHeader(t *testing.T) {
}
}
func TestDownloadAttachmentContentUsesExternalRequestClass(t *testing.T) {
platform := signatureRoundTripFunc(func(req *http.Request) (*http.Response, error) {
return &http.Response{
StatusCode: http.StatusBadGateway,
Header: make(http.Header),
Body: http.NoBody,
Request: req,
}, nil
})
external := signatureRoundTripFunc(func(req *http.Request) (*http.Response, error) {
return &http.Response{
StatusCode: http.StatusOK,
Header: make(http.Header),
Body: io.NopCloser(strings.NewReader("attachment data")),
Request: req,
}, nil
})
rt := newDownloadRuntime(t, &http.Client{
Transport: internaltransport.NewHTTPPolicyRouter(platform, external),
})
data, err := downloadAttachmentContent(rt, "https://open.feishu.cn/presigned/file")
if err != nil {
t.Fatalf("downloadAttachmentContent() error = %v, want external route", err)
}
if got := string(data); got != "attachment data" {
t.Fatalf("downloadAttachmentContent() = %q, want external payload", got)
}
}
// ---------------------------------------------------------------------------
// newOutputRuntime — helper for tests that call runtime.Out / runtime.IO()
// ---------------------------------------------------------------------------

View File

@@ -228,7 +228,7 @@ func downloadSignatureImage(runtime *common.RuntimeContext, downloadURL, filenam
return nil, "", mailInvalidResponseError("signature image download: URL has no host")
}
httpClient, err := runtime.Factory.HttpClient()
httpClient, err := runtime.Factory.ExternalHTTPClient()
if err != nil {
return nil, "", errs.NewInternalError(errs.SubtypeSDKError, "signature image download: %v", err).WithCause(err)
}

View File

@@ -6,15 +6,19 @@ package mail
import (
"context"
"errors"
"io"
"net/http"
"strings"
"testing"
"github.com/spf13/cobra"
"github.com/larksuite/cli/errs"
"github.com/larksuite/cli/internal/auth"
"github.com/larksuite/cli/internal/cmdutil"
"github.com/larksuite/cli/internal/core"
"github.com/larksuite/cli/internal/httpmock"
internaltransport "github.com/larksuite/cli/internal/transport"
"github.com/larksuite/cli/shortcuts/common"
draftpkg "github.com/larksuite/cli/shortcuts/mail/draft"
"github.com/larksuite/cli/shortcuts/mail/emlbuilder"
@@ -115,6 +119,47 @@ func TestContentTypeFromFilename(t *testing.T) {
}
}
func TestDownloadSignatureImageUsesExternalRequestClass(t *testing.T) {
const payload = `{"code":21000,"msg":"application-defined response","data":{"cli_hint":"external"}}`
platform := signatureRoundTripFunc(func(req *http.Request) (*http.Response, error) {
return &http.Response{
StatusCode: http.StatusBadGateway,
Header: make(http.Header),
Body: http.NoBody,
Request: req,
}, nil
})
external := signatureRoundTripFunc(func(req *http.Request) (*http.Response, error) {
return &http.Response{
StatusCode: http.StatusOK,
Header: http.Header{"Content-Type": []string{"application/json"}},
Body: io.NopCloser(strings.NewReader(payload)),
Request: req,
}, nil
})
client := &http.Client{Transport: internaltransport.NewHTTPPolicyRouter(
&auth.SecurityPolicyTransport{Base: platform},
external,
)}
factory := &cmdutil.Factory{
HttpClient: func() (*http.Client, error) { return client, nil },
}
runtime := common.TestNewRuntimeContextWithCtx(context.Background(), &cobra.Command{}, nil)
runtime.Factory = factory
data, contentType, err := downloadSignatureImage(
runtime,
"https://open.feishu.cn/signature.png",
"signature.png",
)
if err != nil {
t.Fatalf("downloadSignatureImage() error = %v, want external response passthrough", err)
}
if string(data) != payload || contentType != "application/json" {
t.Fatalf("download = %q (%s), want external response", data, contentType)
}
}
func TestSignatureCIDsNilSig(t *testing.T) {
if cids := signatureCIDs(nil); cids != nil {
t.Fatalf("expected nil slice for nil sig, got %v", cids)
@@ -206,6 +251,12 @@ func stubSigListResponse(reg *httpmock.Registry, mailboxID string, sigs []map[st
})
}
type signatureRoundTripFunc func(*http.Request) (*http.Response, error)
func (f signatureRoundTripFunc) RoundTrip(req *http.Request) (*http.Response, error) {
return f(req)
}
func TestAutoResolveSignatureID_APIFailureReturnsEmpty(t *testing.T) {
rt, reg := newSigTestRuntime(t)
reg.Register(&httpmock.Stub{

View File

@@ -145,13 +145,8 @@ var MinutesDownload = common.Shortcut{
seen := make(map[string]int)
usedNames := make(map[string]bool)
// Clone the factory client for download use. We clone the struct (not the
// pointer) to avoid mutating the shared singleton's Timeout. The original
// transport chain is preserved so security headers and test mocks still work.
// SSRF protection: ValidateDownloadSourceURL (URL-level) + CheckRedirect
// (redirect-level). Transport-level IP check is intentionally omitted because
// download URLs originate from the trusted Lark API, not user input.
baseClient, err := runtime.Factory.HttpClient()
// Clone the external client so timeout changes stay local.
baseClient, err := runtime.Factory.ExternalHTTPClient()
if err != nil {
return errs.NewNetworkError(errs.SubtypeNetworkTransport, "failed to get HTTP client: %s", err).WithCause(err)
}

View File

@@ -8,6 +8,7 @@ import (
"context"
"encoding/json"
"errors"
"io"
"net/http"
"os"
"strings"
@@ -21,6 +22,7 @@ import (
"github.com/larksuite/cli/internal/core"
"github.com/larksuite/cli/internal/httpmock"
"github.com/larksuite/cli/internal/output"
internaltransport "github.com/larksuite/cli/internal/transport"
"github.com/larksuite/cli/shortcuts/common"
)
@@ -30,6 +32,12 @@ import (
var warmOnce sync.Once
type minutesRoundTripFunc func(*http.Request) (*http.Response, error)
func (f minutesRoundTripFunc) RoundTrip(req *http.Request) (*http.Response, error) {
return f(req)
}
func warmTokenCache(t *testing.T) {
t.Helper()
warmOnce.Do(func() {
@@ -216,6 +224,52 @@ func TestDownload_ServerFilenameTraversalStaysInOutputDir(t *testing.T) {
}
}
func TestDownloadUsesExternalRequestClass(t *testing.T) {
chdir(t, t.TempDir())
const downloadURL = "https://open.feishu.cn/presigned/download"
f, _, _, reg := cmdutil.TestFactory(t, defaultConfig())
reg.Register(mediaStub("tok001", downloadURL))
platform := minutesRoundTripFunc(func(req *http.Request) (*http.Response, error) {
return &http.Response{
StatusCode: http.StatusBadGateway,
Header: make(http.Header),
Body: http.NoBody,
Request: req,
}, nil
})
external := minutesRoundTripFunc(func(req *http.Request) (*http.Response, error) {
payload := []byte("media")
return &http.Response{
StatusCode: http.StatusOK,
Header: http.Header{"Content-Type": []string{"video/mp4"}},
Body: io.NopCloser(bytes.NewReader(payload)),
ContentLength: int64(len(payload)),
Request: req,
}, nil
})
f.HttpClient = func() (*http.Client, error) {
return &http.Client{
Transport: internaltransport.NewHTTPPolicyRouter(platform, external),
}, nil
}
err := mountAndRun(t, MinutesDownload, []string{
"+download", "--minute-tokens", "tok001", "--output", "out.mp4", "--as", "bot",
}, f, nil)
if err != nil {
t.Fatalf("MinutesDownload error = %v, want external route", err)
}
data, err := os.ReadFile("out.mp4")
if err != nil {
t.Fatal(err)
}
if got := string(data); got != "media" {
t.Fatalf("downloaded content = %q, want external payload", got)
}
}
func TestResolveFilenameFromResponse_EmptyDispositionFilename(t *testing.T) {
resp := &http.Response{
Header: http.Header{
@@ -742,7 +796,7 @@ func TestDownload_TypedErr_NetworkTransport_HttpError(t *testing.T) {
Command: "+probe-dl",
AuthTypes: []string{"bot"},
Execute: func(ctx context.Context, rctx *common.RuntimeContext) error {
client, err := rctx.Factory.HttpClient()
client, err := rctx.Factory.ExternalHTTPClient()
if err != nil {
return err
}
@@ -845,7 +899,7 @@ func TestDownload_TypedErr_OverwriteProtection(t *testing.T) {
Command: "+probe-overwrite",
AuthTypes: []string{"bot"},
Execute: func(ctx context.Context, rctx *common.RuntimeContext) error {
client, err := rctx.Factory.HttpClient()
client, err := rctx.Factory.ExternalHTTPClient()
if err != nil {
return err
}

File diff suppressed because it is too large Load Diff

View File

@@ -1004,17 +1004,9 @@ class XmlTextOverlapLintGeometryTest(unittest.TestCase):
"""
)
overlap_pairs = {tuple(issue["elements"]) for issue in result["slides"][0]["issues"]}
# Two caption/label pairs overlap; blV also trips the width-wrap rule (its 15-char
# caption renders wider than its 150px box), which is an intentional error here.
self.assertEqual(result["summary"]["error_count"], 3)
self.assertEqual(result["summary"]["error_count"], 2)
self.assertIn(("blY", "blV"), overlap_pairs)
self.assertIn(("blQ", "blS"), overlap_pairs)
wrap_ids = {
issue["elements"][0]
for issue in result["slides"][0]["issues"]
if issue.get("overflow_axis") == "width"
}
self.assertEqual(wrap_ids, {"blV"})
def test_lint_xml_detects_horizontal_text_overflow_across_declared_box_gap(self) -> None:
result = xml_text_overlap_lint.lint_xml(
@@ -1127,42 +1119,6 @@ class XmlTextOverlapLintGeometryTest(unittest.TestCase):
self.assertEqual(overflowing_issue["overflow"], 30)
self.assertIn('wrap="true" autoFit="normal-auto-fit"', overflowing_issue["message"])
def test_lint_xml_detects_short_label_that_wraps_by_width(self) -> None:
result = xml_text_overlap_lint.lint_xml(
"""
<slide xmlns="http://www.larkoffice.com/sml/2.0">
<data>
<shape id="near-fit" type="text" topLeftX="40" topLeftY="40" width="176" height="96">
<content textType="sub-headline" fontSize="32" fontFamily="思源黑体" bold="true"><p>Slides 87% </p></content>
</shape>
<shape id="under-measured" type="text" topLeftX="40" topLeftY="160" width="136" height="90">
<content fontSize="30" fontFamily="黑体"><p>Docs 99%</p></content>
</shape>
<shape id="auto-fit-spaced" type="text" topLeftX="300" topLeftY="40" width="227" height="96">
<content textType="sub-headline" fontSize="32" fontFamily="思源黑体" bold="true" autoFit="shape-auto-fit"><p>autofix 87% </p></content>
</shape>
<shape id="no-wrap-label" type="text" topLeftX="300" topLeftY="160" width="136" height="90">
<content fontSize="30" fontFamily="黑体" wrap="false"><p>Docs 99%</p></content>
</shape>
<shape id="comfortable" type="text" topLeftX="600" topLeftY="40" width="300" height="60">
<content fontSize="24" fontFamily="思源黑体"><p>OK</p></content>
</shape>
</data>
</slide>
"""
)
wrap_issues = [
issue for issue in result["slides"][0]["issues"] if issue.get("overflow_axis") == "width"
]
wrap_ids = {issue["elements"][0] for issue in wrap_issues}
# The three real false-negatives are caught, independent of autoFit and collapsed spaces.
self.assertEqual(wrap_ids, {"near-fit", "under-measured", "auto-fit-spaced"})
self.assertTrue(all(issue["level"] == "error" for issue in wrap_issues))
self.assertTrue(all(issue["code"] == "text_may_overflow_shape" for issue in wrap_issues))
# wrap="false" opts a run out; a label that comfortably fits is not flagged.
self.assertNotIn("no-wrap-label", wrap_ids)
self.assertNotIn("comfortable", wrap_ids)
def test_lint_xml_uses_fixed_line_spacing_for_text_height_warning(self) -> None:
result = xml_text_overlap_lint.lint_xml(
"""
@@ -1247,28 +1203,6 @@ class XmlTextOverlapLintGeometryTest(unittest.TestCase):
]
self.assertEqual(overflow_issues, [])
def test_lint_xml_reports_labeled_short_metric_when_it_wraps(self) -> None:
result = xml_text_overlap_lint.lint_xml(
"""
<slide xmlns="http://www.larkoffice.com/sml/2.0">
<data>
<shape id="sheet-success" type="text" topLeftX="520" topLeftY="385" width="180" height="50">
<content textType="headline" fontSize="32" bold="true" autoFit="no-auto-fit">
<p>Sheet 98.5%</p>
</content>
</shape>
</data>
</slide>
"""
)
overflow_issues = [
issue
for issue in result["slides"][0]["issues"]
if issue["code"] == "text_may_overflow_shape"
]
self.assertEqual(len(overflow_issues), 1)
self.assertEqual(overflow_issues[0]["elements"], ["sheet-success"])
def test_lint_xml_reports_plain_short_metric_when_it_wraps(self) -> None:
result = xml_text_overlap_lint.lint_xml(
"""
@@ -1289,97 +1223,6 @@ class XmlTextOverlapLintGeometryTest(unittest.TestCase):
self.assertEqual(len(overflow_issues), 1)
self.assertEqual(overflow_issues[0]["elements"], ["plain-age"])
def test_lint_xml_reports_cjk_credit_with_em_dashes_wrapping_narrow_box(self) -> None:
# "—— 李白" in a tight author-credit box wraps in the renderer because the two em-dashes render
# full-width inside a CJK run (slides p1: bMW). unicodedata marks em-dash as ambiguous width, so
# a naive Latin-punctuation estimate under-reports the line and misses the wrap. The width check
# must treat ambiguous glyphs as full-width in CJK context (Bucket A4).
result = xml_text_overlap_lint.lint_xml(
"""
<slide xmlns="http://www.larkoffice.com/sml/2.0">
<data>
<shape id="credit" type="text" topLeftX="66" topLeftY="124" width="46" height="18">
<content fontSize="12"><p>—— 李白</p></content>
</shape>
</data>
</slide>
"""
)
wrap_issues = [
issue for issue in result["slides"][0]["issues"]
if issue["code"] == "text_may_overflow_shape" and issue["elements"] == ["credit"]
]
# Promoting the em-dashes to full-width makes the run too wide for the box; the renderer then
# wraps it to two lines that also overflow the 18px height, so either the width or the height
# detector may surface it first -- the contract is that the credit is flagged, not which axis.
self.assertEqual(len(wrap_issues), 1)
self.assertIn(wrap_issues[0]["overflow_axis"], {"width", "height"})
def test_lint_xml_keeps_latin_en_dash_range_narrow(self) -> None:
# The ambiguous-width promotion is context-gated: an en-dash in a pure-Latin run ("20202023")
# stays half-width, so a comfortably-sized box must not be reported. Guards A4 from over-firing
# by inflating every dash to full-width regardless of surrounding script.
result = xml_text_overlap_lint.lint_xml(
"""
<slide xmlns="http://www.larkoffice.com/sml/2.0">
<data>
<shape id="range" type="text" topLeftX="80" topLeftY="80" width="140" height="30">
<content fontSize="14"><p>20202023</p></content>
</shape>
</data>
</slide>
"""
)
wrap_issues = [
issue for issue in result["slides"][0]["issues"]
if issue["code"] == "text_may_overflow_shape"
]
self.assertEqual(wrap_issues, [])
def test_lint_xml_reports_percent_heavy_run_overflowing_by_full_width_glyph(self) -> None:
# "%" is Unicode half-width (Na) but renders near full-width, so a percentage-heavy run wraps to
# more lines than a naive punct-coefficient estimate and overflows its box height (slides p3:
# bhU "Docs 99%Docs 99%Docs 99%%1"). Measuring "%" at its true advance is what surfaces this.
result = xml_text_overlap_lint.lint_xml(
"""
<slide xmlns="http://www.larkoffice.com/sml/2.0">
<data>
<shape id="metrics" type="text" topLeftX="587" topLeftY="60" width="227" height="100">
<content fontSize="30" autoFit="no-auto-fit"><p>Docs 99%Docs 99%Docs 99%%1</p></content>
</shape>
</data>
</slide>
"""
)
overflow = [
issue for issue in result["slides"][0]["errors"]
if issue["code"] == "text_may_overflow_shape" and issue["elements"] == ["metrics"]
]
self.assertEqual(len(overflow), 1)
def test_lint_xml_marginal_height_warning_does_not_mask_width_error(self) -> None:
# A short "Slides 87%" label sized so its wrapped two lines graze the box height by <1px yields a
# height *warning*, while the same run is genuinely too wide -> a width *error*. The width error
# must still surface: a marginal height warning must not suppress it via already_flagged_ids
# (slides p3: bMP/bhd). The run is reported once, at error level.
result = xml_text_overlap_lint.lint_xml(
"""
<slide xmlns="http://www.larkoffice.com/sml/2.0">
<data>
<shape id="label" type="text" topLeftX="244" topLeftY="120" width="176" height="79">
<content fontSize="32" bold="true" autoFit="no-auto-fit"><p>Slides 87%</p></content>
</shape>
</data>
</slide>
"""
)
reports = [
issue for issue in result["slides"][0]["issues"]
if issue["code"] == "text_may_overflow_shape" and issue["elements"] == ["label"]
]
self.assertEqual(len(reports), 1)
self.assertEqual(reports[0]["level"], "error")
def test_lint_xml_allows_centered_short_label_near_fit_as_single_line(self) -> None:
result = xml_text_overlap_lint.lint_xml(
"""
@@ -1873,10 +1716,10 @@ class XmlTextOverlapLintGeometryTest(unittest.TestCase):
<slide xmlns="http://www.larkoffice.com/sml/2.0">
<data>
<img src="tok" topLeftX="-120" topLeftY="20" width="360" height="360"/>
<shape type="text" topLeftX="300" topLeftY="80" width="180" height="80">
<shape type="text" topLeftX="40" topLeftY="80" width="180" height="80">
<content textType="title" fontSize="44"><p>Title</p></content>
</shape>
<shape type="text" topLeftX="300" topLeftY="170" width="180" height="40">
<shape type="text" topLeftX="40" topLeftY="120" width="180" height="40">
<content textType="sub-headline" fontSize="20"><p>Subtitle</p></content>
</shape>
</data>
@@ -2050,32 +1893,6 @@ class XmlTextOverlapLintGeometryTest(unittest.TestCase):
]
self.assertEqual(len(crossing), 1)
def test_lint_xml_reports_horizontal_line_inside_wide_line_spacing_span(self) -> None:
# 3 lines of fontSize 20 at multiple:1.8 give a real 92px glyph span, but the flat
# font_size*1.2 approximation is only 72px. Both boxes centre in the 200px shape, so the flat
# eroded box is ~[246,314] while the spacing-aware eroded box is ~[236,324]. A rule at y=240
# lands in that top margin -- inside the real glyph rows yet outside the flat box -- so it only
# reports once the line-crossing path uses the spacing-aware height (Bucket C, slides p8/p10).
result = xml_text_overlap_lint.lint_xml(
"""
<slide xmlns="http://www.larkoffice.com/sml/2.0">
<data>
<shape id="poem" type="text" topLeftX="80" topLeftY="180" width="360" height="200">
<content fontSize="20" lineSpacing="multiple:1.8"><p>第一行诗句文字</p><p>第二行诗句文字</p><p>第三行诗句文字</p></content>
</shape>
<line id="rule" startX="80" startY="240" endX="220" endY="240">
<border color="rgb(0, 0, 0)" width="3"/>
</line>
</data>
</slide>
"""
)
crossing = [
issue for issue in result["slides"][0]["errors"] if set(issue["elements"]) == {"rule", "poem"}
]
self.assertEqual(len(crossing), 1)
self.assertEqual(crossing[0]["code"], "bbox_overlap")
def test_lint_xml_reports_diagonal_line_crossing_text_block(self) -> None:
result = xml_text_overlap_lint.lint_xml(
"""
@@ -2504,266 +2321,6 @@ class XmlTextOverlapLintGeometryTest(unittest.TestCase):
self.assertEqual(result["summary"]["warning_count"], 0)
self.assertEqual(result["slides"][0]["issues"], [])
def test_lint_xml_reports_rotated_text_colliding_with_horizontal_text(self) -> None:
# A 270-rotated label sweeps a vertical footprint that overlaps a nearby horizontal label. With
# rotation-aware glyph boxes the collision is detectable, and because the runs are not parallel
# the overlap ratio is tiny so the absolute-area fallback must flag it (slides p6, Bucket D+E).
result = xml_text_overlap_lint.lint_xml(
"""
<slide xmlns="http://www.larkoffice.com/sml/2.0">
<data>
<shape id="flat" type="text" topLeftX="240" topLeftY="200" width="64" height="24">
<content fontSize="16"><p>文字碰撞</p></content>
</shape>
<shape id="spun" type="text" topLeftX="272" topLeftY="232" width="64" height="24" rotation="270">
<content fontSize="16"><p>文字碰撞</p></content>
</shape>
</data>
</slide>
"""
)
collisions = [
issue for issue in result["slides"][0]["errors"]
if issue["code"] == "bbox_overlap" and set(issue["elements"]) == {"flat", "spun"}
]
self.assertEqual(len(collisions), 1)
def test_lint_xml_still_suppresses_coincident_shadow_text_overlay(self) -> None:
# A drop-shadow duplicate offset by a pixel is an intentional overlay; the coincidence check
# must keep suppressing it even though the text is identical (guards the E1 tightening).
result = xml_text_overlap_lint.lint_xml(
"""
<slide xmlns="http://www.larkoffice.com/sml/2.0">
<data>
<shape id="shadow" type="text" topLeftX="200" topLeftY="200" width="200" height="40">
<content fontSize="20"><p>标题文字</p></content>
</shape>
<shape id="fill" type="text" topLeftX="202" topLeftY="202" width="200" height="40">
<content fontSize="20"><p>标题文字</p></content>
</shape>
</data>
</slide>
"""
)
collisions = [
issue for issue in result["slides"][0]["errors"]
if issue["code"] == "bbox_overlap" and set(issue["elements"]) == {"shadow", "fill"}
]
self.assertEqual(collisions, [])
def test_lint_xml_reports_text_overflowing_background_container(self) -> None:
# Text anchored inside a background card whose glyph box spills past the card's bottom edge has
# outgrown the box the author sized for it (slides p7). The card is drawn first (lower z-order),
# so it is the container; the text must surface as text_overflows_container (Bucket B).
result = xml_text_overlap_lint.lint_xml(
"""
<slide xmlns="http://www.larkoffice.com/sml/2.0">
<data>
<shape id="card" type="rect" topLeftX="200" topLeftY="200" width="120" height="40">
<fill><fillColor color="rgba(230,230,230,1)"/></fill>
</shape>
<shape id="body" type="text" topLeftX="205" topLeftY="205" width="110" height="120">
<content fontSize="16"><p>第一行</p><p>第二行</p><p>第三行</p></content>
</shape>
</data>
</slide>
"""
)
overflow = [
issue for issue in result["slides"][0]["errors"]
if issue["code"] == "text_overflows_container" and set(issue["elements"]) == {"body", "card"}
]
self.assertEqual(len(overflow), 1)
self.assertGreater(overflow[0]["overflow"]["bottom"], 4)
def test_lint_xml_ignores_text_fitting_inside_background_container(self) -> None:
# Text whose glyph box stays inside its background card is fine; the container rule must stay
# silent so tightly-fitted-but-valid cards are not falsely reported.
result = xml_text_overlap_lint.lint_xml(
"""
<slide xmlns="http://www.larkoffice.com/sml/2.0">
<data>
<shape id="card" type="rect" topLeftX="200" topLeftY="200" width="200" height="120">
<fill><fillColor color="rgba(230,230,230,1)"/></fill>
</shape>
<shape id="body" type="text" topLeftX="210" topLeftY="210" width="180" height="40">
<content fontSize="14"><p>短文本</p></content>
</shape>
</data>
</slide>
"""
)
codes = [issue["code"] for issue in result["slides"][0]["issues"]]
self.assertNotIn("text_overflows_container", codes)
def test_lint_xml_reports_free_text_shape_overlapping_table_grid(self) -> None:
# A free-floating text shape whose glyph box lands on top of a sibling table occludes the cell
# contents (slides p4). The table renders its own text; a stray shape over the grid is an
# accidental overlay, so it must surface as table_covers_text (Bucket B).
result = xml_text_overlap_lint.lint_xml(
"""
<presentation xmlns="http://www.larkoffice.com/sml/2.0" width="960" height="540">
<slide xmlns="http://www.larkoffice.com/sml/2.0">
<data>
<table id="grid" topLeftX="200" topLeftY="200" width="400" height="150">
<tr><td><content><p>A</p></content></td></tr>
</table>
<shape id="stray" type="text" topLeftX="260" topLeftY="240" width="120" height="30">
<content fontSize="16"><p>覆盖表格</p></content>
</shape>
</data>
</slide>
</presentation>
"""
)
occlusions = [
issue for issue in result["slides"][0]["errors"]
if issue["code"] == "table_covers_text" and set(issue["elements"]) == {"grid", "stray"}
]
self.assertEqual(len(occlusions), 1)
def test_lint_xml_ignores_table_with_only_cell_text(self) -> None:
# Cell text is part of the table's own layout and is never extracted as a standalone shape, so
# a table alone must not self-report table_covers_text (guards against a runaway detector).
result = xml_text_overlap_lint.lint_xml(
"""
<presentation xmlns="http://www.larkoffice.com/sml/2.0" width="960" height="540">
<slide xmlns="http://www.larkoffice.com/sml/2.0">
<data>
<table id="solo" topLeftX="200" topLeftY="200" width="400" height="150">
<tr><td><content><p>Score</p></content></td></tr>
</table>
</data>
</slide>
</presentation>
"""
)
codes = [issue["code"] for issue in result["slides"][0]["issues"]]
self.assertNotIn("table_covers_text", codes)
def test_lint_xml_reports_free_text_shape_overlapping_chart(self) -> None:
# A free-floating text shape whose glyph box lands on top of a sibling chart occludes the chart's
# generated labels and legend (slides p5: a headline dropped onto a pie chart's ring). The chart
# renders its own text; a stray shape over the plot area is an accidental overlay, so it must
# surface as chart_covers_text (Bucket B3).
result = xml_text_overlap_lint.lint_xml(
"""
<presentation xmlns="http://www.larkoffice.com/sml/2.0" width="960" height="540">
<slide xmlns="http://www.larkoffice.com/sml/2.0">
<data>
<chart id="pie" topLeftX="200" topLeftY="60" width="420" height="420">
<chartData><dim1><chartField name="p">A,B</chartField></dim1></chartData>
</chart>
<shape id="stray" type="text" topLeftX="360" topLeftY="120" width="120" height="40">
<content fontSize="32"><p>abc 99%</p></content>
</shape>
</data>
</slide>
</presentation>
"""
)
occlusions = [
issue for issue in result["slides"][0]["errors"]
if issue["code"] == "chart_covers_text" and set(issue["elements"]) == {"pie", "stray"}
]
self.assertEqual(len(occlusions), 1)
def test_lint_xml_ignores_chart_not_overlapping_text(self) -> None:
# A chart and a text shape that sit side by side without their glyph boxes touching must not
# report chart_covers_text (guards the detector from firing on mere co-existence).
result = xml_text_overlap_lint.lint_xml(
"""
<presentation xmlns="http://www.larkoffice.com/sml/2.0" width="960" height="540">
<slide xmlns="http://www.larkoffice.com/sml/2.0">
<data>
<chart id="pie" topLeftX="40" topLeftY="60" width="300" height="300">
<chartData><dim1><chartField name="p">A,B</chartField></dim1></chartData>
</chart>
<shape id="caption" type="text" topLeftX="600" topLeftY="80" width="200" height="40">
<content fontSize="16"><p>Sales breakdown</p></content>
</shape>
</data>
</slide>
</presentation>
"""
)
codes = [issue["code"] for issue in result["slides"][0]["issues"]]
self.assertNotIn("chart_covers_text", codes)
def test_lint_xml_reports_auto_fit_title_growing_onto_body_below(self) -> None:
# A shape-auto-fit title sized for one line wraps to two, growing downward past its authored box
# onto the body text beneath it (slides p9). shape-auto-fit only means the box grows to fit, so
# the grown glyph height -- not the authored height -- is what collides. The body is a tall
# multi-line block so the overlap covers <30% of it: the generic text-text check cannot catch
# this, only the dedicated auto-fit growth detector can (Bucket A3 / auto-fit growth).
result = xml_text_overlap_lint.lint_xml(
"""
<slide xmlns="http://www.larkoffice.com/sml/2.0">
<data>
<shape id="title" type="text" topLeftX="80" topLeftY="20" width="480" height="36">
<content fontSize="24" autoFit="shape-auto-fit"><p>02. | Literature Review - International Research</p></content>
</shape>
<shape id="body" type="text" topLeftX="80" topLeftY="60" width="480" height="200">
<content fontSize="15" verticalAlign="top"><p>1. Marxist Perspective</p><p>Line two of body copy</p><p>Line three of body copy</p><p>Line four of body copy</p><p>Line five of body copy</p><p>Line six of body copy</p></content>
</shape>
</data>
</slide>
"""
)
collisions = [
issue for issue in result["slides"][0]["errors"]
if issue["code"] == "bbox_overlap" and set(issue["elements"]) == {"title", "body"}
]
self.assertEqual(len(collisions), 1)
def test_lint_xml_ignores_auto_fit_title_with_space_below(self) -> None:
# An identical wrapping auto-fit title with an empty gap below it grows harmlessly; the check
# must stay silent so ordinary auto-fit growth is not flagged (guards the grown-region area gate).
result = xml_text_overlap_lint.lint_xml(
"""
<slide xmlns="http://www.larkoffice.com/sml/2.0">
<data>
<shape id="title" type="text" topLeftX="80" topLeftY="20" width="480" height="36">
<content fontSize="24" autoFit="shape-auto-fit"><p>02. | Literature Review - International Research</p></content>
</shape>
<shape id="body" type="text" topLeftX="80" topLeftY="300" width="480" height="200">
<content fontSize="15"><p>1. Marxist Perspective</p></content>
</shape>
</data>
</slide>
"""
)
collisions = [
issue for issue in result["slides"][0]["issues"]
if issue["code"] == "bbox_overlap" and set(issue["elements"]) == {"title", "body"}
]
self.assertEqual(collisions, [])
def test_lint_xml_does_not_treat_divider_rule_as_text_background_container(self) -> None:
# A thin horizontal rule under a title is a divider, not a container. Owning a title's grown
# glyph box to a 3px rule and reporting it as text_overflows_container is a false positive
# (slides p9); the line-like guard must keep the divider out of the container candidate set.
result = xml_text_overlap_lint.lint_xml(
"""
<slide xmlns="http://www.larkoffice.com/sml/2.0">
<data>
<shape id="rule" type="rect" topLeftX="40" topLeftY="60" width="880" height="3">
<fill><fillColor color="rgba(40,60,120,1)"/></fill>
</shape>
<shape id="title" type="text" topLeftX="80" topLeftY="20" width="480" height="36">
<content fontSize="24" autoFit="shape-auto-fit"><p>02. | Literature Review - International Research</p></content>
</shape>
</data>
</slide>
"""
)
container_hits = [
issue for issue in result["slides"][0]["issues"]
if issue["code"] == "text_overflows_container" and "rule" in issue["elements"]
]
self.assertEqual(container_hits, [])
def test_lint_xml_keeps_resolved_table_sizes_positive_when_target_is_too_small(self) -> None:
result = xml_text_overlap_lint.lint_xml(
"""
@@ -2902,79 +2459,6 @@ class XmlTextOverlapLintGeometryTest(unittest.TestCase):
self.assertEqual(result["summary"]["error_count"], 0)
self.assertEqual(result["summary"]["info_count"], 1)
def test_lint_xml_reports_image_text_overlap_even_when_image_precedes_text_in_xml_order(self) -> None:
result = xml_text_overlap_lint.lint_xml(
"""
<slide xmlns="http://www.larkoffice.com/sml/2.0"><data>
<img id="image" src="token" topLeftX="120" topLeftY="120" width="120" height="60"/>
<shape id="text" type="text" topLeftX="100" topLeftY="100" width="220" height="90">
<content fontSize="28" lineSpacing="fixed:34" wrap="false"><p>Quarterly Plan</p></content>
</shape>
</data></slide>
"""
)
issue = next(issue for issue in result["slides"][0]["issues"] if issue["code"] == "image_covers_text")
self.assertEqual(issue["elements"], ["image", "text"])
self.assertIn("no longer overlaps the text glyph area", issue["hint"])
self.assertEqual(result["summary"]["error_count"], 1)
def test_lint_xml_exempts_full_canvas_background_image_behind_text(self) -> None:
# A full-bleed image at the bottom of the z-order is the slide backdrop; text rendered on top of
# it is never occluded (slides p9: bBo fills the whole canvas under the content). It must not be
# reported as image_covers_text (Bucket B5 background-image false positive).
result = xml_text_overlap_lint.lint_xml(
"""
<presentation xmlns="http://www.larkoffice.com/sml/2.0" width="960" height="540">
<slide xmlns="http://www.larkoffice.com/sml/2.0"><data>
<img id="backdrop" src="token" topLeftX="0" topLeftY="0" width="960" height="540"/>
<shape id="text" type="text" topLeftX="100" topLeftY="100" width="400" height="60">
<content fontSize="28"><p>On the backdrop</p></content>
</shape>
</data></slide>
</presentation>
"""
)
codes = [
issue for issue in result["slides"][0]["issues"]
if issue["code"] == "image_covers_text" and "backdrop" in issue["elements"]
]
self.assertEqual(codes, [])
def test_lint_xml_reports_full_canvas_image_drawn_above_text(self) -> None:
# The exemption is z-order aware: a full-canvas image drawn *after* (above) the text really does
# cover it, so it must still be flagged. Guards the backdrop exemption from swallowing real
# occlusions where the image is on top.
result = xml_text_overlap_lint.lint_xml(
"""
<presentation xmlns="http://www.larkoffice.com/sml/2.0" width="960" height="540">
<slide xmlns="http://www.larkoffice.com/sml/2.0"><data>
<shape id="text" type="text" topLeftX="100" topLeftY="100" width="400" height="60">
<content fontSize="28"><p>Under the cover</p></content>
</shape>
<img id="cover" src="token" topLeftX="0" topLeftY="0" width="960" height="540"/>
</data></slide>
</presentation>
"""
)
codes = [
issue for issue in result["slides"][0]["issues"]
if issue["code"] == "image_covers_text" and set(issue["elements"]) == {"cover", "text"}
]
self.assertEqual(len(codes), 1)
def test_stacking_helpers_agree_on_paint_order(self) -> None:
lower = {"order": 1}
upper = {"order": 3}
same = {"order": 3}
# is_drawn_behind and is_drawn_in_front_of are strict and mutually exclusive inverses.
self.assertTrue(xml_text_overlap_lint.is_drawn_behind(lower, upper))
self.assertFalse(xml_text_overlap_lint.is_drawn_in_front_of(lower, upper))
self.assertTrue(xml_text_overlap_lint.is_drawn_in_front_of(upper, lower))
self.assertFalse(xml_text_overlap_lint.is_drawn_behind(upper, lower))
# Equal order is neither behind nor in front, so an equal-order sibling never occludes.
self.assertFalse(xml_text_overlap_lint.is_drawn_behind(same, upper))
self.assertFalse(xml_text_overlap_lint.is_drawn_in_front_of(same, upper))
class XmlTextOverlapLintDensityTest(unittest.TestCase):
def test_lint_xml_blocks_blank_slide(self) -> None: