mirror of
https://github.com/infiniflow/ragflow.git
synced 2026-08-12 11:43:39 +08:00
Go: add context to lots of interface (#17253)
Signed-off-by: Jin Hai <haijin.chn@gmail.com>
This commit is contained in:
@@ -784,7 +784,7 @@ func TestBotRoutes_RequireAuth(t *testing.T) {
|
||||
func TestBotMiddleware_NonBearerRegularToken(t *testing.T) {
|
||||
gin.SetMode(gin.TestMode)
|
||||
stub := &stubUserTokenResolver{
|
||||
getUserByTokenFn: func(auth string) (*entity.User, common.ErrorCode, error) {
|
||||
getUserByTokenFn: func(ctx context.Context, auth string) (*entity.User, common.ErrorCode, error) {
|
||||
if auth != "raw-access-token-abc" {
|
||||
t.Errorf("GetUserByToken called with %q, want raw-access-token-abc", auth)
|
||||
}
|
||||
@@ -820,36 +820,36 @@ func TestBotMiddleware_NonBearerRegularToken(t *testing.T) {
|
||||
// return safe defaults (CodeUnauthorized so the middleware
|
||||
// short-circuits to 401).
|
||||
type stubUserTokenResolver struct {
|
||||
getUserByTokenFn func(authorization string) (*entity.User, common.ErrorCode, error)
|
||||
getUserByAPITokenFn func(token string) (*entity.User, common.ErrorCode, error)
|
||||
getUserByBetaAPITokenFn func(token string) (*entity.User, common.ErrorCode, error)
|
||||
getAPITokenByBetaFn func(authorization string) (*entity.APIToken, error)
|
||||
getUserByTokenFn func(ctx context.Context, authorization string) (*entity.User, common.ErrorCode, error)
|
||||
getUserByAPITokenFn func(ctx context.Context, token string) (*entity.User, common.ErrorCode, error)
|
||||
getUserByBetaAPITokenFn func(ctx context.Context, token string) (*entity.User, common.ErrorCode, error)
|
||||
getAPITokenByBetaFn func(ctx context.Context, authorization string) (*entity.APIToken, error)
|
||||
}
|
||||
|
||||
func (s *stubUserTokenResolver) GetUserByToken(authorization string) (*entity.User, common.ErrorCode, error) {
|
||||
func (s *stubUserTokenResolver) GetUserByToken(ctx context.Context, authorization string) (*entity.User, common.ErrorCode, error) {
|
||||
if s.getUserByTokenFn != nil {
|
||||
return s.getUserByTokenFn(authorization)
|
||||
return s.getUserByTokenFn(ctx, authorization)
|
||||
}
|
||||
return nil, common.CodeUnauthorized, errors.New("not stubbed")
|
||||
}
|
||||
|
||||
func (s *stubUserTokenResolver) GetUserByAPIToken(token string) (*entity.User, common.ErrorCode, error) {
|
||||
func (s *stubUserTokenResolver) GetUserByAPIToken(ctx context.Context, token string) (*entity.User, common.ErrorCode, error) {
|
||||
if s.getUserByAPITokenFn != nil {
|
||||
return s.getUserByAPITokenFn(token)
|
||||
return s.getUserByAPITokenFn(ctx, token)
|
||||
}
|
||||
return nil, common.CodeUnauthorized, errors.New("not stubbed")
|
||||
}
|
||||
|
||||
func (s *stubUserTokenResolver) GetUserByBetaAPIToken(token string) (*entity.User, common.ErrorCode, error) {
|
||||
func (s *stubUserTokenResolver) GetUserByBetaAPIToken(ctx context.Context, token string) (*entity.User, common.ErrorCode, error) {
|
||||
if s.getUserByBetaAPITokenFn != nil {
|
||||
return s.getUserByBetaAPITokenFn(token)
|
||||
return s.getUserByBetaAPITokenFn(ctx, token)
|
||||
}
|
||||
return nil, common.CodeUnauthorized, errors.New("not stubbed")
|
||||
}
|
||||
|
||||
func (s *stubUserTokenResolver) GetAPITokenByBeta(authorization string) (*entity.APIToken, error) {
|
||||
func (s *stubUserTokenResolver) GetAPITokenByBeta(ctx context.Context, authorization string) (*entity.APIToken, error) {
|
||||
if s.getAPITokenByBetaFn != nil {
|
||||
return s.getAPITokenByBetaFn(authorization)
|
||||
return s.getAPITokenByBetaFn(ctx, authorization)
|
||||
}
|
||||
return nil, errors.New("not stubbed")
|
||||
}
|
||||
@@ -870,7 +870,7 @@ func (s *stubUserTokenResolver) GetAPITokenByBeta(authorization string) (*entity
|
||||
func TestBotRoutes_NoRegularAuthRequired(t *testing.T) {
|
||||
gin.SetMode(gin.TestMode)
|
||||
stub := &stubUserTokenResolver{
|
||||
getUserByTokenFn: func(auth string) (*entity.User, common.ErrorCode, error) {
|
||||
getUserByTokenFn: func(ctx context.Context, auth string) (*entity.User, common.ErrorCode, error) {
|
||||
return &entity.User{ID: "u-regular"}, common.CodeSuccess, nil
|
||||
},
|
||||
}
|
||||
@@ -955,7 +955,7 @@ func TestBotRoutes_NoRegularAuthRequired(t *testing.T) {
|
||||
func TestDownloadAttachment_Unauth(t *testing.T) {
|
||||
gin.SetMode(gin.TestMode)
|
||||
stub := &stubUserTokenResolver{
|
||||
getUserByTokenFn: func(auth string) (*entity.User, common.ErrorCode, error) {
|
||||
getUserByTokenFn: func(ctx context.Context, auth string) (*entity.User, common.ErrorCode, error) {
|
||||
return nil, common.CodeUnauthorized, errors.New("invalid token")
|
||||
},
|
||||
}
|
||||
@@ -968,13 +968,14 @@ func TestDownloadAttachment_Unauth(t *testing.T) {
|
||||
// resolve via GetUserByToken. Both branches must reject
|
||||
// with 401.
|
||||
g.Use(func(c *gin.Context) {
|
||||
ctx := c.Request.Context()
|
||||
auth := c.GetHeader("Authorization")
|
||||
if auth == "" {
|
||||
common.ResponseWithCodeData(c, common.CodeUnauthorized, nil, "Authorization required")
|
||||
c.Abort()
|
||||
return
|
||||
}
|
||||
if u, code, err := stub.GetUserByToken(auth); err != nil || code != common.CodeSuccess {
|
||||
if u, code, err := stub.GetUserByToken(ctx, auth); err != nil || code != common.CodeSuccess {
|
||||
common.ResponseWithCodeData(c, common.CodeUnauthorized, nil, "Invalid auth credentials")
|
||||
c.Abort()
|
||||
return
|
||||
|
||||
Reference in New Issue
Block a user