Add wsl linter and fix

Signed-off-by: Émile Ré <emile@probo.com>
This commit is contained in:
Émile Ré
2026-05-19 14:51:08 +04:00
parent eedfdcecc8
commit 9156d6a16a
882 changed files with 6068 additions and 574 deletions

View File

@@ -42,11 +42,13 @@ func (c *APIKeyConnection) Client(ctx context.Context) (*http.Client, error) {
tokenType: "Bearer",
underlying: httpclient.DefaultPooledTransport(httpclient.WithSSRFProtection()),
}
return &http.Client{Transport: transport}, nil
}
func (c APIKeyConnection) MarshalJSON() ([]byte, error) {
type Alias APIKeyConnection
return json.Marshal(&struct {
Type string `json:"type"`
Alias
@@ -58,10 +60,12 @@ func (c APIKeyConnection) MarshalJSON() ([]byte, error) {
func (c *APIKeyConnection) UnmarshalJSON(data []byte) error {
type Alias APIKeyConnection
aux := &struct {
*Alias
}{
Alias: (*Alias)(c),
}
return json.Unmarshal(data, &aux)
}

View File

@@ -70,6 +70,7 @@ func UnmarshalConnection(protocol string, provider string, data []byte) (Connect
if err := json.Unmarshal(data, &slackConn); err != nil {
return nil, fmt.Errorf("cannot unmarshal slack connection: %w", err)
}
return &slackConn, nil
default:
@@ -77,6 +78,7 @@ func UnmarshalConnection(protocol string, provider string, data []byte) (Connect
if err := json.Unmarshal(data, &conn); err != nil {
return nil, fmt.Errorf("cannot unmarshal oauth2 connection: %w", err)
}
return &conn, nil
}
@@ -85,6 +87,7 @@ func UnmarshalConnection(protocol string, provider string, data []byte) (Connect
if err := json.Unmarshal(data, &conn); err != nil {
return nil, fmt.Errorf("cannot unmarshal api key connection: %w", err)
}
return &conn, nil
}

View File

@@ -143,11 +143,13 @@ func (c *OAuth2Connector) Initiate(
ConnectorID: opts.ConnectorID,
RequestedScopes: opts.Scopes,
}
if r != nil {
if continueURL := r.URL.Query().Get("continue"); continueURL != "" {
stateData.ContinueURL = continueURL
}
}
return c.InitiateWithState(ctx, stateData, opts)
}
@@ -166,6 +168,7 @@ func (c *OAuth2Connector) InitiateWithState(
if err != nil {
return "", fmt.Errorf("cannot generate PKCE verifier: %w", err)
}
stateData.CodeVerifier = verifier
}
@@ -179,6 +182,7 @@ func (c *OAuth2Connector) InitiateWithState(
authCodeQuery.Set("client_id", c.ClientID)
authCodeQuery.Set("redirect_uri", c.RedirectURI)
authCodeQuery.Set("response_type", "code")
if len(opts.Scopes) > 0 {
authCodeQuery.Set("scope", strings.Join(opts.Scopes, " "))
}
@@ -200,6 +204,7 @@ func (c *OAuth2Connector) InitiateWithState(
if incrementalAuth && k == "prompt" && v == "consent" {
continue
}
authCodeQuery.Set(k, v)
}
@@ -220,6 +225,7 @@ func generatePKCEVerifier() (string, error) {
if _, err := rand.Read(b); err != nil {
return "", fmt.Errorf("cannot read random bytes: %w", err)
}
return base64.RawURLEncoding.EncodeToString(b), nil
}
@@ -277,6 +283,7 @@ func (c *OAuth2Connector) CompleteWithState(ctx context.Context, r *http.Request
if err != nil {
return nil, nil, fmt.Errorf("cannot post token URL: %w", err)
}
defer func() { _ = tokenResp.Body.Close() }()
if tokenResp.StatusCode != http.StatusOK {
@@ -351,6 +358,7 @@ func (c *OAuth2Connector) buildTokenRequest(ctx context.Context, code, redirectU
if codeVerifier != "" {
body["code_verifier"] = codeVerifier
}
jsonBody, err := json.Marshal(body)
if err != nil {
return nil, fmt.Errorf("cannot marshal token request body: %w", err)
@@ -370,6 +378,7 @@ func (c *OAuth2Connector) buildTokenRequest(ctx context.Context, code, redirectU
req.Header.Set("Accept", "application/json")
req.Header.Set("User-Agent", "Probo Connector")
req.Header.Set("Authorization", basicAuthHeader(c.ClientID, c.ClientSecret))
return req, nil
case "basic-form":
@@ -378,6 +387,7 @@ func (c *OAuth2Connector) buildTokenRequest(ctx context.Context, code, redirectU
formData.Set("code", code)
formData.Set("redirect_uri", redirectURI)
formData.Set("grant_type", "authorization_code")
if codeVerifier != "" {
formData.Set("code_verifier", codeVerifier)
}
@@ -396,6 +406,7 @@ func (c *OAuth2Connector) buildTokenRequest(ctx context.Context, code, redirectU
req.Header.Set("Accept", "application/json")
req.Header.Set("User-Agent", "Probo Connector")
req.Header.Set("Authorization", basicAuthHeader(c.ClientID, c.ClientSecret))
return req, nil
default:
@@ -406,6 +417,7 @@ func (c *OAuth2Connector) buildTokenRequest(ctx context.Context, code, redirectU
formData.Set("code", code)
formData.Set("redirect_uri", redirectURI)
formData.Set("grant_type", "authorization_code")
if codeVerifier != "" {
formData.Set("code_verifier", codeVerifier)
}
@@ -423,6 +435,7 @@ func (c *OAuth2Connector) buildTokenRequest(ctx context.Context, code, redirectU
req.Header.Set("Content-Type", "application/x-www-form-urlencoded; charset=utf-8")
req.Header.Set("Accept", "application/json")
req.Header.Set("User-Agent", "Probo Connector")
return req, nil
}
}
@@ -457,6 +470,7 @@ func (c *OAuth2Connection) ClientWithOptions(ctx context.Context, opts ...httpcl
client := &http.Client{
Transport: transport,
}
return client, nil
}
@@ -481,6 +495,7 @@ func (c *OAuth2Connection) RefreshableClient(ctx context.Context, cfg OAuth2Refr
// Determine auth style based on TokenEndpointAuth
authStyle := oauth2.AuthStyleInParams
switch cfg.TokenEndpointAuth {
case "basic-form", "basic-json":
authStyle = oauth2.AuthStyleInHeader
@@ -529,6 +544,7 @@ func (c *OAuth2Connection) RefreshableClient(ctx context.Context, cfg OAuth2Refr
// Update the connection with the potentially refreshed token
c.AccessToken = newToken.AccessToken
c.ExpiresAt = newToken.Expiry
c.TokenType = newToken.TokenType
if newToken.RefreshToken != "" {
c.RefreshToken = newToken.RefreshToken
@@ -558,6 +574,7 @@ func (c *OAuth2Connection) clientCredentialsClient(ctx context.Context, opts ...
formData := url.Values{}
formData.Set("grant_type", "client_credentials")
if c.Scope != "" {
formData.Set("scope", c.Scope)
}
@@ -585,6 +602,7 @@ func (c *OAuth2Connection) clientCredentialsClient(ctx context.Context, opts ...
if err != nil {
return nil, fmt.Errorf("cannot post client credentials token URL: %w", err)
}
defer func() { _ = resp.Body.Close() }()
if resp.StatusCode != http.StatusOK {
@@ -609,9 +627,11 @@ func (c *OAuth2Connection) clientCredentialsClient(ctx context.Context, opts ...
if rawToken.TokenType != "" {
c.TokenType = rawToken.TokenType
}
if c.TokenType == "" {
c.TokenType = "Bearer"
}
if rawToken.ExpiresIn > 0 {
c.ExpiresAt = time.Now().Add(time.Duration(rawToken.ExpiresIn) * time.Second)
}
@@ -627,6 +647,7 @@ func (c *OAuth2Connection) clientCredentialsClient(ctx context.Context, opts ...
func (c OAuth2Connection) MarshalJSON() ([]byte, error) {
type Alias OAuth2Connection
return json.Marshal(&struct {
Type string `json:"type"`
Alias
@@ -638,11 +659,13 @@ func (c OAuth2Connection) MarshalJSON() ([]byte, error) {
func (c *OAuth2Connection) UnmarshalJSON(data []byte) error {
type Alias OAuth2Connection
aux := &struct {
*Alias
}{
Alias: (*Alias)(c),
}
return json.Unmarshal(data, &aux)
}
@@ -660,5 +683,6 @@ func (t *oauth2Transport) RoundTrip(req *http.Request) (*http.Response, error) {
// string), so we always send "Bearer" -- the only scheme any connector in
// this codebase actually needs.
req2.Header.Set("Authorization", "Bearer "+t.token)
return t.underlying.RoundTrip(req2)
}

View File

@@ -26,5 +26,6 @@ func (g OAuth2GrantType) IsValid() bool {
case OAuth2GrantTypeAuthorizationCode, OAuth2GrantTypeClientCredentials:
return true
}
return false
}

View File

@@ -183,6 +183,7 @@ func TestBuildTokenRequest_BasicJSON(t *testing.T) {
require.NoError(t, err)
var jsonBody map[string]string
err = json.Unmarshal(body, &jsonBody)
require.NoError(t, err)
@@ -193,6 +194,7 @@ func TestBuildTokenRequest_BasicJSON(t *testing.T) {
// JSON body must NOT contain client credentials
_, hasClientID := jsonBody["client_id"]
_, hasClientSecret := jsonBody["client_secret"]
assert.False(t, hasClientID, "JSON body should not contain client_id")
assert.False(t, hasClientSecret, "JSON body should not contain client_secret")
}
@@ -636,12 +638,14 @@ func TestInitiateWithState_PKCE(t *testing.T) {
t.Parallel()
var capturedVerifier string
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
body, err := io.ReadAll(r.Body)
assert.NoError(t, err)
form, err := url.ParseQuery(string(body))
assert.NoError(t, err)
capturedVerifier = form.Get("code_verifier")
w.Header().Set("Content-Type", "application/json")
@@ -676,6 +680,7 @@ func TestInitiateWithState_PKCE(t *testing.T) {
payload, err := DecodeOAuth2StatePayload(stateToken)
require.NoError(t, err)
expectedVerifier := payload.Data.CodeVerifier
require.NotEmpty(t, expectedVerifier)
@@ -705,11 +710,13 @@ func TestApplyProviderDefaults_AuthURLTemplating(t *testing.T) {
// test so we do not have to wait for a real Vercel-style provider
// to land. Restore on teardown.
const fakeProvider = "TEST_TEMPLATED_AUTH_URL"
previous, hadPrevious := providerDefinitions[fakeProvider]
providerDefinitions[fakeProvider] = providerDefinition{
AuthURL: "https://example.com/integrations/{integration_slug}/new",
TokenURL: "https://example.com/oauth/token",
}
t.Cleanup(func() {
if hadPrevious {
providerDefinitions[fakeProvider] = previous

View File

@@ -29,14 +29,17 @@ func AbsorbPagerDutyTokenResponse(state *OAuth2State, body []byte) {
if state == nil || state.Provider != PagerDutyProvider {
return
}
var pd struct {
Subdomain string `json:"subdomain"`
}
if err := json.Unmarshal(body, &pd); err != nil || pd.Subdomain == "" {
return
}
if state.ProviderMetadata == nil {
state.ProviderMetadata = map[string]string{}
}
state.ProviderMetadata["subdomain"] = pd.Subdomain
}

View File

@@ -39,21 +39,25 @@ func NewConnectorRegistry() *ConnectorRegistry {
func (r *ConnectorRegistry) Register(provider string, c Connector) error {
r.Lock()
defer r.Unlock()
if _, ok := r.connectors[provider]; ok {
return fmt.Errorf("cannot register connector %q: already registered", provider)
}
r.connectors[provider] = c
return nil
}
func (r *ConnectorRegistry) Get(provider string) (Connector, error) {
r.RLock()
defer r.RUnlock()
c, ok := r.connectors[provider]
if !ok {
return nil, fmt.Errorf("cannot find connector %q", provider)
}
return c, nil
}

View File

@@ -31,16 +31,21 @@ func ParseScopeString(s string) []string {
if len(fields) == 0 {
return []string{}
}
seen := make(map[string]struct{}, len(fields))
out := make([]string, 0, len(fields))
for _, f := range fields {
if _, ok := seen[f]; ok {
continue
}
seen[f] = struct{}{}
out = append(out, f)
}
sort.Strings(out)
return out
}
@@ -50,9 +55,11 @@ func FormatScopeString(scopes []string) string {
if len(scopes) == 0 {
return ""
}
sorted := make([]string, len(scopes))
copy(sorted, scopes)
sort.Strings(sorted)
return strings.Join(sorted, " ")
}
@@ -61,18 +68,23 @@ func FormatScopeString(scopes []string) string {
// result is a fresh slice and never aliases any input.
func UnionScopes(scopeSets ...[]string) []string {
seen := map[string]struct{}{}
for _, set := range scopeSets {
for _, s := range set {
if s == "" {
continue
}
seen[s] = struct{}{}
}
}
out := make([]string, 0, len(seen))
for s := range seen {
out = append(out, s)
}
sort.Strings(out)
return out
}

View File

@@ -120,9 +120,11 @@ func ParseSlackTokenResponse(body []byte, oauth2Conn OAuth2Connection, organizat
if slackResponse.Error != "" {
return nil, nil, fmt.Errorf("cannot complete Slack OAuth2 flow: %s", slackResponse.Error)
}
if !slackResponse.Ok {
return nil, nil, fmt.Errorf("cannot complete Slack OAuth2 flow: ok=false")
}
if oauth2Conn.AccessToken == "" {
return nil, nil, fmt.Errorf("cannot complete Slack OAuth2 flow: missing access token")
}

View File

@@ -36,6 +36,7 @@ func TestParseSlackTokenResponse(t *testing.T) {
t.Run("with incoming webhook", func(t *testing.T) {
t.Parallel()
body := []byte(`{"ok":true,"incoming_webhook":{"url":"https://hooks.slack.com/services/T/B/X","channel":"#general","channel_id":"C123"}}`)
conn, returnedOrgID, err := ParseSlackTokenResponse(body, base, orgID)
@@ -52,6 +53,7 @@ func TestParseSlackTokenResponse(t *testing.T) {
t.Run("without incoming webhook", func(t *testing.T) {
t.Parallel()
body := []byte(`{"ok":true}`)
conn, returnedOrgID, err := ParseSlackTokenResponse(body, base, orgID)
@@ -67,6 +69,7 @@ func TestParseSlackTokenResponse(t *testing.T) {
t.Run("slack error response", func(t *testing.T) {
t.Parallel()
body := []byte(`{"ok":false,"error":"invalid_code"}`)
conn, returnedOrgID, err := ParseSlackTokenResponse(body, base, orgID)
@@ -78,6 +81,7 @@ func TestParseSlackTokenResponse(t *testing.T) {
t.Run("missing access token", func(t *testing.T) {
t.Parallel()
body := []byte(`{"ok":true}`)
connWithoutToken := base
connWithoutToken.AccessToken = ""

View File

@@ -43,12 +43,14 @@ func FetchVercelUser(ctx context.Context, client *http.Client) (VercelUser, erro
if err != nil {
return VercelUser{}, fmt.Errorf("cannot create vercel user request: %w", err)
}
req.Header.Set("Accept", "application/json")
resp, err := client.Do(req)
if err != nil {
return VercelUser{}, fmt.Errorf("cannot execute vercel user request: %w", err)
}
defer func() { _ = resp.Body.Close() }()
if resp.StatusCode < 200 || resp.StatusCode >= 300 {
@@ -61,6 +63,7 @@ func FetchVercelUser(ctx context.Context, client *http.Client) (VercelUser, erro
if err := json.NewDecoder(resp.Body).Decode(&body); err != nil {
return VercelUser{}, fmt.Errorf("cannot decode vercel user response: %w", err)
}
return body.User, nil
}
@@ -77,14 +80,17 @@ func FetchVercelUserID(ctx context.Context, accessToken string) (string, error)
if err != nil {
return "", fmt.Errorf("cannot create vercel user request: %w", err)
}
req.Header.Set("Accept", "application/json")
req.Header.Set("Authorization", "Bearer "+accessToken)
client := httpclient.DefaultClient(httpclient.WithSSRFProtection())
resp, err := client.Do(req)
if err != nil {
return "", fmt.Errorf("cannot execute vercel user request: %w", err)
}
defer func() { _ = resp.Body.Close() }()
if resp.StatusCode < 200 || resp.StatusCode >= 300 {
@@ -97,5 +103,6 @@ func FetchVercelUserID(ctx context.Context, accessToken string) (string, error)
if err := json.NewDecoder(resp.Body).Decode(&body); err != nil {
return "", fmt.Errorf("cannot decode vercel user response: %w", err)
}
return body.User.ID, nil
}