Skip to content

Commit a81f63f

Browse files
committed
refactor: Extract sso proxy auth to own middleware
1 parent a77e965 commit a81f63f

5 files changed

Lines changed: 212 additions & 214 deletions

File tree

internal/http/middleware/auth.go

Lines changed: 3 additions & 72 deletions
Original file line numberDiff line numberDiff line change
@@ -1,9 +1,7 @@
11
package middleware
22

33
import (
4-
"errors"
54
"fmt"
6-
"net"
75
"net/http"
86
"strings"
97

@@ -16,40 +14,14 @@ import (
1614
// request context.
1715
type AuthMiddleware struct {
1816
deps model.Dependencies
19-
20-
trustedIPs []*net.IPNet
2117
}
2218

2319
func NewAuthMiddleware(deps model.Dependencies) *AuthMiddleware {
24-
plainIPs := deps.Config().Http.SSOProxyAuthTrusted
25-
trustedIPs := make([]*net.IPNet, len(plainIPs))
26-
for i, ip := range plainIPs {
27-
_, ipNet, err := net.ParseCIDR(ip)
28-
if err != nil {
29-
deps.Logger().WithError(err).WithField("ip", ip).Error("Failed to parse trusted ip cidr")
30-
continue
31-
}
32-
33-
trustedIPs[i] = ipNet
34-
}
35-
36-
return &AuthMiddleware{
37-
deps: deps,
38-
trustedIPs: trustedIPs,
39-
}
20+
return &AuthMiddleware{deps: deps}
4021
}
4122

4223
func (m *AuthMiddleware) OnRequest(deps model.Dependencies, c model.WebContext) error {
43-
account, err := m.ssoAccount(deps, c)
44-
if err != nil {
45-
deps.Logger().
46-
WithError(err).
47-
WithField("remote_addr", c.Request().RemoteAddr).
48-
WithField("request_id", c.GetRequestID()).
49-
Error("getting sso account")
50-
}
51-
if account != nil {
52-
c.SetAccount(account)
24+
if c.UserIsLogged() {
5325
return nil
5426
}
5527

@@ -62,7 +34,7 @@ func (m *AuthMiddleware) OnRequest(deps model.Dependencies, c model.WebContext)
6234
return nil
6335
}
6436

65-
account, err = deps.Domains().Auth().CheckToken(c.Request().Context(), token)
37+
account, err := deps.Domains().Auth().CheckToken(c.Request().Context(), token)
6638
if err != nil {
6739
// If we fail to check token, remove the token cookie and redirect to login
6840
deps.Logger().WithError(err).WithField("request_id", c.GetRequestID()).Error("Failed to check token")
@@ -78,47 +50,6 @@ func (m *AuthMiddleware) OnRequest(deps model.Dependencies, c model.WebContext)
7850
return nil
7951
}
8052

81-
func (m *AuthMiddleware) ssoAccount(deps model.Dependencies, c model.WebContext) (*model.AccountDTO,error) {
82-
if !deps.Config().Http.SSOProxyAuth {
83-
return nil, nil
84-
}
85-
86-
remoteAddr := c.Request().RemoteAddr
87-
ip, _, err := net.SplitHostPort(remoteAddr)
88-
if err != nil {
89-
var addrErr *net.AddrError
90-
if errors.As(err, &addrErr) && addrErr.Err == "missing port in address" {
91-
ip = remoteAddr
92-
} else {
93-
return nil,err
94-
}
95-
}
96-
requestIP := net.ParseIP(ip)
97-
if !m.isTrustedIP(requestIP) {
98-
return nil, errors.New("remoteAddr is not a trusted ip")
99-
}
100-
101-
headerName := deps.Config().Http.SSOProxyAuthHeaderName
102-
userName := c.Request().Header.Get(headerName)
103-
if userName == "" {
104-
return nil, nil
105-
}
106-
107-
account, err := deps.Domains().Accounts().GetAccountByUsername(c.Request().Context(), userName)
108-
if err != nil {
109-
return nil, err
110-
}
111-
112-
return account, nil
113-
}
114-
func (m *AuthMiddleware) isTrustedIP(ip net.IP) bool {
115-
for _, net := range m.trustedIPs {
116-
if ok := net.Contains(ip); ok {
117-
return true
118-
}
119-
}
120-
return false
121-
}
12253

12354
func (m *AuthMiddleware) OnResponse(deps model.Dependencies, c model.WebContext) error {
12455
return nil
Lines changed: 104 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,104 @@
1+
package middleware
2+
3+
import (
4+
"errors"
5+
"net"
6+
7+
"github.com/go-shiori/shiori/internal/model"
8+
)
9+
10+
// AuthMiddleware handles authentication for incoming request by checking the token
11+
// from the Authorization header or the token cookie and setting the account in the
12+
// request context.
13+
type AuthSSOProxyMiddleware struct {
14+
deps model.Dependencies
15+
16+
trustedIPs []*net.IPNet
17+
}
18+
19+
func NewAuthSSOProxyMiddleware(deps model.Dependencies) *AuthSSOProxyMiddleware {
20+
plainIPs := deps.Config().Http.SSOProxyAuthTrusted
21+
trustedIPs := make([]*net.IPNet, len(plainIPs))
22+
for i, ip := range plainIPs {
23+
_, ipNet, err := net.ParseCIDR(ip)
24+
if err != nil {
25+
deps.Logger().WithError(err).WithField("ip", ip).Error("Failed to parse trusted ip cidr")
26+
continue
27+
}
28+
29+
trustedIPs[i] = ipNet
30+
}
31+
32+
return &AuthSSOProxyMiddleware{
33+
deps: deps,
34+
trustedIPs: trustedIPs,
35+
}
36+
}
37+
38+
func (m *AuthSSOProxyMiddleware) OnRequest(deps model.Dependencies, c model.WebContext) error {
39+
if c.UserIsLogged() {
40+
return nil
41+
}
42+
43+
account, err := m.ssoAccount(deps, c)
44+
if err != nil {
45+
deps.Logger().
46+
WithError(err).
47+
WithField("remote_addr", c.Request().RemoteAddr).
48+
WithField("request_id", c.GetRequestID()).
49+
Error("getting sso account")
50+
return nil
51+
}
52+
if account != nil {
53+
c.SetAccount(account)
54+
return nil
55+
}
56+
57+
return nil
58+
}
59+
60+
func (m *AuthSSOProxyMiddleware) ssoAccount(deps model.Dependencies, c model.WebContext) (*model.AccountDTO,error) {
61+
if !deps.Config().Http.SSOProxyAuth {
62+
return nil, nil
63+
}
64+
65+
remoteAddr := c.Request().RemoteAddr
66+
ip, _, err := net.SplitHostPort(remoteAddr)
67+
if err != nil {
68+
var addrErr *net.AddrError
69+
if errors.As(err, &addrErr) && addrErr.Err == "missing port in address" {
70+
ip = remoteAddr
71+
} else {
72+
return nil,err
73+
}
74+
}
75+
requestIP := net.ParseIP(ip)
76+
if !m.isTrustedIP(requestIP) {
77+
return nil, errors.New("remoteAddr is not a trusted ip")
78+
}
79+
80+
headerName := deps.Config().Http.SSOProxyAuthHeaderName
81+
userName := c.Request().Header.Get(headerName)
82+
if userName == "" {
83+
return nil, nil
84+
}
85+
86+
account, err := deps.Domains().Accounts().GetAccountByUsername(c.Request().Context(), userName)
87+
if err != nil {
88+
return nil, err
89+
}
90+
91+
return account, nil
92+
}
93+
func (m *AuthSSOProxyMiddleware) isTrustedIP(ip net.IP) bool {
94+
for _, net := range m.trustedIPs {
95+
if ok := net.Contains(ip); ok {
96+
return true
97+
}
98+
}
99+
return false
100+
}
101+
102+
func (m *AuthSSOProxyMiddleware) OnResponse(deps model.Dependencies, c model.WebContext) error {
103+
return nil
104+
}
Lines changed: 101 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,101 @@
1+
package middleware
2+
3+
import (
4+
"context"
5+
"net/http"
6+
"net/http/httptest"
7+
"testing"
8+
9+
"github.com/go-shiori/shiori/internal/http/webcontext"
10+
"github.com/go-shiori/shiori/internal/model"
11+
"github.com/go-shiori/shiori/internal/testutil"
12+
"github.com/sirupsen/logrus"
13+
"github.com/stretchr/testify/require"
14+
)
15+
16+
func TestAuthMiddlewareWithSSO(t *testing.T) {
17+
logger := logrus.New()
18+
_, deps := testutil.GetTestConfigurationAndDependencies(t, context.TODO(), logger)
19+
deps.Config().Http.SSOProxyAuth = true
20+
21+
account, err := deps.Domains().Accounts().CreateAccount(context.TODO(), model.AccountDTO{
22+
ID: model.DBID(98),
23+
Username: "test_username",
24+
Password: "super_secure_password",
25+
})
26+
require.NoError(t, err)
27+
28+
t.Run("test no authorization method", func(t *testing.T) {
29+
w := httptest.NewRecorder()
30+
r := httptest.NewRequest(http.MethodGet, "/", nil)
31+
c := webcontext.NewWebContext(w, r)
32+
33+
middleware := NewAuthSSOProxyMiddleware(deps)
34+
err := middleware.OnRequest(deps, c)
35+
require.NoError(t, err)
36+
require.Nil(t, c.GetAccount())
37+
})
38+
39+
t.Run("test untrusted ip", func(t *testing.T) {
40+
w := httptest.NewRecorder()
41+
r := httptest.NewRequest(http.MethodGet, "/", nil)
42+
r.RemoteAddr = "invalid-ip"
43+
c := webcontext.NewWebContext(w, r)
44+
45+
middleware := NewAuthSSOProxyMiddleware(deps)
46+
err := middleware.OnRequest(deps, c)
47+
require.NoError(t, err)
48+
require.Nil(t, c.GetAccount())
49+
})
50+
51+
t.Run("test empty header", func(t *testing.T) {
52+
w := httptest.NewRecorder()
53+
r := httptest.NewRequest(http.MethodGet, "/", nil)
54+
r.RemoteAddr = "10.0.0.3"
55+
c := webcontext.NewWebContext(w, r)
56+
57+
middleware := NewAuthSSOProxyMiddleware(deps)
58+
err := middleware.OnRequest(deps, c)
59+
require.NoError(t, err)
60+
require.Nil(t, c.GetAccount())
61+
})
62+
63+
t.Run("test invalid sso username", func(t *testing.T) {
64+
w := httptest.NewRecorder()
65+
r := httptest.NewRequest(http.MethodGet, "/", nil)
66+
r.RemoteAddr = "10.0.0.3"
67+
r.Header.Add("Remote-User", "username")
68+
c := webcontext.NewWebContext(w, r)
69+
70+
middleware := NewAuthSSOProxyMiddleware(deps)
71+
err := middleware.OnRequest(deps, c)
72+
require.NoError(t, err)
73+
require.Nil(t, c.GetAccount())
74+
})
75+
76+
t.Run("test sso login", func(t *testing.T) {
77+
w := httptest.NewRecorder()
78+
r := httptest.NewRequest(http.MethodGet, "/", nil)
79+
r.RemoteAddr = "10.0.0.3"
80+
r.Header.Add("Remote-User", account.Username)
81+
c := webcontext.NewWebContext(w, r)
82+
83+
middleware := NewAuthSSOProxyMiddleware(deps)
84+
err := middleware.OnRequest(deps, c)
85+
require.NoError(t, err)
86+
require.NotNil(t, c.GetAccount())
87+
})
88+
89+
t.Run("test sso login ip:port", func(t *testing.T) {
90+
w := httptest.NewRecorder()
91+
r := httptest.NewRequest(http.MethodGet, "/", nil)
92+
r.RemoteAddr = "10.0.0.3:65342"
93+
r.Header.Add("Remote-User", account.Username)
94+
c := webcontext.NewWebContext(w, r)
95+
96+
middleware := NewAuthSSOProxyMiddleware(deps)
97+
err := middleware.OnRequest(deps, c)
98+
require.NoError(t, err)
99+
require.NotNil(t, c.GetAccount())
100+
})
101+
}

0 commit comments

Comments
 (0)