diff --git a/services/proxy/pkg/middleware/authentication.go b/services/proxy/pkg/middleware/authentication.go
index 4a4e441f5d..b5dc7d5326 100644
--- a/services/proxy/pkg/middleware/authentication.go
+++ b/services/proxy/pkg/middleware/authentication.go
@@ -337,18 +337,17 @@ func renderTerminalFailure(w http.ResponseWriter, r *http.Request, result Authen
}
// Default: generic 401 without challenges.
- w.WriteHeader(http.StatusUnauthorized)
- if webdav.IsWebdavRequest(r) {
- b, err := webdav.Marshal(webdav.Exception{
- Code: webdav.SabredavNotAuthenticated,
- Message: "Authentication error",
- })
- if err != nil {
- return err
- }
- _, _ = w.Write(b)
+ if !webdav.IsWebdavRequest(r) {
+ w.WriteHeader(http.StatusUnauthorized)
+ return nil
}
- return nil
+
+ w.Header().Set("Content-Type", "application/xml; charset=utf-8")
+ w.WriteHeader(http.StatusUnauthorized)
+ return webdav.Encode(w, webdav.Exception{
+ Code: webdav.SabredavNotAuthenticated,
+ Message: "Authentication error",
+ })
}
// renderErrorDetails renders a response based on structured error details.
@@ -377,20 +376,16 @@ func renderJSONExpired(w http.ResponseWriter, permissionID string) error {
PermissionID: permissionID,
}
- body, err := json.Marshal(resp)
- if err != nil {
- return err
- }
-
w.Header().Set("Content-Type", "application/json")
w.WriteHeader(http.StatusUnauthorized)
- _, err = w.Write(body)
- return err
+ return json.NewEncoder(w).Encode(resp)
}
// renderDAVExpired writes a SabreDAV-compatible XML response for expired guest sessions.
func renderDAVExpired(w http.ResponseWriter, permissionID string) error {
- xmlBody, err := xml.Marshal(davExpiredResponse{
+ w.Header().Set("Content-Type", "application/xml; charset=utf-8")
+ w.WriteHeader(http.StatusUnauthorized)
+ return webdav.EncodeXML(w, davExpiredResponse{
XmlnsD: "DAV",
XmlnsS: "http://sabredav.org/ns",
Exception: "Sabre\\DAV\\Exception\\NotAuthenticated",
@@ -401,14 +396,6 @@ func renderDAVExpired(w http.ResponseWriter, permissionID string) error {
ShareID: permissionID,
},
})
- if err != nil {
- return err
- }
-
- w.Header().Set("Content-Type", "application/xml; charset=utf-8")
- w.WriteHeader(http.StatusUnauthorized)
- _, err = w.Write(append([]byte(xml.Header), xmlBody...))
- return err
}
type davExpiredResponse struct {
diff --git a/services/proxy/pkg/middleware/terminal_failure_test.go b/services/proxy/pkg/middleware/terminal_failure_test.go
new file mode 100644
index 0000000000..8039d0c119
--- /dev/null
+++ b/services/proxy/pkg/middleware/terminal_failure_test.go
@@ -0,0 +1,87 @@
+package middleware
+
+import (
+ "encoding/json"
+ "encoding/xml"
+ "net/http"
+ "net/http/httptest"
+ "strings"
+ "testing"
+
+ "github.com/stretchr/testify/assert"
+ "github.com/stretchr/testify/require"
+)
+
+// hostileShareID would break out of an XML/JSON document if it were not escaped.
+const hostileShareID = `&'`
+
+func TestRenderTerminalFailure_Generic(t *testing.T) {
+ t.Run("non-dav request", func(t *testing.T) {
+ rr := httptest.NewRecorder()
+ req := httptest.NewRequest(http.MethodGet, "/graph/v1beta1/me/drive/sharedWithMe", nil)
+
+ require.NoError(t, renderTerminalFailure(rr, req, TerminalFailed()))
+
+ assert.Equal(t, http.StatusUnauthorized, rr.Code)
+ assert.Empty(t, rr.Body.String())
+ })
+
+ t.Run("dav request", func(t *testing.T) {
+ rr := httptest.NewRecorder()
+ req := httptest.NewRequest("PROPFIND", "/dav/spaces/abc", nil)
+
+ require.NoError(t, renderTerminalFailure(rr, req, TerminalFailed()))
+
+ assert.Equal(t, http.StatusUnauthorized, rr.Code)
+ assert.Equal(t, "application/xml; charset=utf-8", rr.Header().Get("Content-Type"))
+ assert.True(t, strings.HasPrefix(rr.Body.String(), xml.Header))
+ assert.Contains(t, rr.Body.String(), "Sabre\\DAV\\Exception\\NotAuthenticated")
+ })
+}
+
+func TestRenderTerminalFailure_GuestSessionExpiredJSON(t *testing.T) {
+ rr := httptest.NewRecorder()
+ req := httptest.NewRequest(http.MethodGet, "/graph/v1beta1/me/drive/sharedWithMe", nil)
+
+ require.NoError(t, renderTerminalFailure(rr, req, AuthenticationResult{
+ State: AuthenticationFailed,
+ Terminal: true,
+ ErrorDetails: GuestSessionExpiredDetails{PermissionID: hostileShareID},
+ }))
+
+ assert.Equal(t, http.StatusUnauthorized, rr.Code)
+ assert.Equal(t, "application/json", rr.Header().Get("Content-Type"))
+
+ var body struct {
+ ErrorType string `json:"error_type"`
+ PermissionID string `json:"permissionId"`
+ }
+ require.NoError(t, json.Unmarshal(rr.Body.Bytes(), &body))
+ assert.Equal(t, "session_expired", body.ErrorType)
+ assert.Equal(t, hostileShareID, body.PermissionID)
+ // json.Encoder escapes <, > and & so the body is safe even if it were rendered as HTML
+ assert.NotContains(t, rr.Body.String(), "