fix: fix objectnode get object range 0-

Signed-off-by: justice <justice_103@126.com>
This commit is contained in:
justice 2021-12-24 17:14:33 +08:00 committed by Shuoran Liu
parent a4d15b37d1
commit 10ebbd6760
3 changed files with 162 additions and 41 deletions

View File

@ -201,6 +201,20 @@ class S3TestCase(TestCase):
body = result['Body'].read()
self.assertEqual(compute_md5(body), body_md5)
def assert_get_object_range_result(self, result, status_code=206, content_length=None, content_range=None):
self.assertNotEqual(result, None)
self.assertEqual(type(result), dict)
self.assertTrue('ResponseMetadata' in result)
self.assertTrue('HTTPStatusCode' in result['ResponseMetadata'])
self.assertEqual(result['ResponseMetadata']['HTTPStatusCode'], status_code)
if status_code == 200:
self.assertEqual(result['ResponseMetadata']['HTTPHeaders']['accept-ranges'], 'bytes')
if content_range is not None:
self.assertEqual(result['ResponseMetadata']['HTTPHeaders']['content-range'], content_range)
if content_length is not None:
self.assertEqual(result['ResponseMetadata']['HTTPHeaders']['content-length'], str(content_length))
def assert_put_object_result(self, result, etag=None, content_type=None, content_length=None):
self.assertNotEqual(result, None)
self.assertEqual(type(result), dict)

View File

@ -0,0 +1,87 @@
# Copyright 2020 The ChubaoFS Authors.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
# http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or
# implied. See the License for the specific language governing
# permissions and limitations under the License.
# -*- coding: utf-8 -*-
from base import S3TestCase
from base import random_string, random_bytes, compute_md5, get_env_s3_client
from env import BUCKET
KEY_PREFIX = 'test-object-get-range/'
class ObjectGetRangeTest(S3TestCase):
'''
'''
s3 = None
def __init__(self, case):
super(ObjectGetRangeTest, self).__init__(case)
self.s3 = get_env_s3_client()
self.file_size = 10000
self.file_key = KEY_PREFIX + random_string(16)
self.test_cases = [
{ "range":"bytes=0-499", "status_code": 206, "content-range":"bytes 0-499/10000", "content-length": 500 },
{ "range":"bytes=500-999", "status_code": 206, "content-range":"bytes 500-999/10000", "content-length": 500 },
{ "range":"bytes=9500-", "status_code": 206, "content-range":"bytes 9500-9999/10000", "content-length": 500 },
{ "range":"bytes=0-", "status_code": 206, "content-range":"bytes 0-9999/10000", "content-length": 10000 },
{ "range":"bytes=0-0", "status_code": 206, "content-range":"bytes 0-0/10000", "content-length": 1 },
{ "range":"bytes=-500", "status_code": 206, "content-range":"bytes 9500-9999/10000", "content-length": 500 },
{ "range":"bytes=-1", "status_code": 206, "content-range":"bytes 9999-9999/10000", "content-length": 1 },
{ "range":"bytes=-0", "status_code": 206, "content-range":"bytes 0-9999/10000", "content-length": 10000 },
{ "range":"bytes=1-0", "status_code": 416 },
{ "range":"bytes=10", "status_code": 416 },
{ "range":"bytes=", "status_code": 416 },
{ "range":"bytes=abc", "status_code": 416 },
{ "range":"bytes=abc-123", "status_code": 416 },
{ "range":"1-0", "status_code": 416 },
]
self._init_object()
def _init_object(self):
file_keys = []
result = self.s3.list_objects(Bucket=BUCKET, Prefix=KEY_PREFIX)
if 'Contents' in result:
contents = result['Contents']
for content in contents:
file_keys.append({'Key': content.get('Key')})
if len(file_keys) > 0:
self.s3.delete_objects(
Bucket=BUCKET,
Delete={'Objects': file_keys}
)
self.s3.put_object(Bucket=BUCKET, Key=self.file_key, Body=random_bytes(self.file_size))
def test_get_object_range(self):
'''
test get object range
'''
for test_case in self.test_cases:
try:
result=self.s3.get_object(Bucket=BUCKET, Key=self.file_key, Range=test_case.get('range')),
except Exception as e:
self.assert_client_error(error=e, expect_status_code=test_case.get('status_code'))
else:
self.assert_get_object_range_result(
result=result[0],
status_code=test_case.get('status_code'),
content_range=test_case.get('content-range', None),
content_length=test_case.get('content-length', None)
)
self.s3.delete_object( Bucket=BUCKET, Key=self.file_key)

View File

@ -33,7 +33,7 @@ import (
)
var (
rangeRegexp = regexp.MustCompile("^bytes=(\\d)+-(\\d)*$")
rangeRegexp = regexp.MustCompile("^bytes=(\\d)*-(\\d)*$")
)
// Get object
@ -75,49 +75,65 @@ func (o *ObjectNode) getObjectHandler(w http.ResponseWriter, r *http.Request) {
var isRangeRead bool
var partSize uint64
var partCount uint64
var lowerPart, upperPart string
if len(rangeOpt) > 0 && rangeRegexp.MatchString(rangeOpt) {
var hyphenIndex = strings.Index(rangeOpt, "-")
if hyphenIndex < 0 {
errorCode = InvalidArgument
revRange := false
if len(rangeOpt) > 0 {
if !rangeRegexp.MatchString(rangeOpt) {
log.LogWarnf("getObjectHandler: getObject fail: requestID(%v) invalid rangeOpt(%v)",
GetRequestID(r), rangeOpt)
errorCode = InvalidRange
return
}
lowerPart = rangeOpt[len("bytes="):hyphenIndex]
upperPart = ""
if hyphenIndex+1 < len(rangeOpt) {
upperPart = rangeOpt[hyphenIndex+1:]
}
if len(lowerPart) > 0 {
if rangeLower, err = strconv.ParseUint(lowerPart, 10, 64); err != nil {
log.LogErrorf("getObjectHandler: parse range lower fail: requestID(%v) rangeOpt(%v) err(%v)",
GetRequestID(r), rangeOpt, err)
ServeInternalStaticErrorResponse(w, r)
} else {
var hyphenIndex = strings.Index(rangeOpt, "-")
if hyphenIndex < 0 {
log.LogWarnf("getObjectHandler: getObject fail: requestID(%v) invalid rangeOpt(%v)",
GetRequestID(r), rangeOpt)
errorCode = InvalidRange
return
}
}
if len(upperPart) > 0 {
if rangeUpper, err = strconv.ParseUint(upperPart, 10, 64); err != nil {
log.LogErrorf("getObjectHandler: parse range upper fail: requestID(%v) rangeOpt(%v) err(%v)",
GetRequestID(r), rangeOpt, err)
ServeInternalStaticErrorResponse(w, r)
return
}
}
if rangeUpper > 0 && rangeUpper < rangeLower {
// upper enabled and lower than lower side
if err = InvalidArgument.ServeResponse(w, r); err != nil {
log.LogErrorf("getObjectHandler: serve response fail: requestID(%v) err(%v)",
GetRequestID(r), err)
return
}
}
isRangeRead = true
log.LogDebugf("getObjectHandler: parse range option: requestID(%v) rangeOpt(%v) rangeLower(%v) rangeUpper(%v)",
GetRequestID(r), rangeOpt, rangeLower, rangeUpper)
var lowerPart = rangeOpt[len("bytes="):hyphenIndex]
var upperPart = ""
if hyphenIndex+1 < len(rangeOpt) {
// bytes=-5
if hyphenIndex == len("bytes=") { //suffix range opt
revRange = true
rangeUpper = 1<<64 - 1
lowerPart = rangeOpt[hyphenIndex+1:]
} else { // bytes=1-10
upperPart = rangeOpt[hyphenIndex+1:]
}
} else if hyphenIndex+1 == len(rangeOpt) { // bytes=1-
rangeUpper = 1<<64 - 1
}
if len(lowerPart) > 0 {
if rangeLower, err = strconv.ParseUint(lowerPart, 10, 64); err != nil {
log.LogErrorf("getObjectHandler: parse range lower fail: requestID(%v) rangeOpt(%v) err(%v)",
GetRequestID(r), rangeOpt, err)
ServeInternalStaticErrorResponse(w, r)
return
}
}
if len(upperPart) > 0 {
if rangeUpper, err = strconv.ParseUint(upperPart, 10, 64); err != nil {
log.LogErrorf("getObjectHandler: parse range upper fail: requestID(%v) rangeOpt(%v) err(%v)",
GetRequestID(r), rangeOpt, err)
ServeInternalStaticErrorResponse(w, r)
return
}
}
if rangeUpper < rangeLower {
// upper enabled and lower than lower side
log.LogWarnf("getObjectHandler: getObject fail: requestID(%v) invalid rangeOpt(%v)",
GetRequestID(r), rangeOpt)
errorCode = InvalidRange
return
}
isRangeRead = true
log.LogDebugf("getObjectHandler: parse range option: requestID(%v) rangeOpt(%v) rangeLower(%v) rangeUpper(%v)",
GetRequestID(r), rangeOpt, rangeLower, rangeUpper)
}
}
responseCacheControl := r.URL.Query().Get(ParamResponseCacheControl)
@ -207,14 +223,18 @@ func (o *ObjectNode) getObjectHandler(w http.ResponseWriter, r *http.Request) {
}
// validate and fix range
if isRangeRead && (rangeUpper > uint64(fileInfo.Size)-1 || upperPart == "") {
if isRangeRead && rangeUpper > uint64(fileInfo.Size)-1 {
rangeUpper = uint64(fileInfo.Size) - 1
if revRange && rangeLower > 0 {
rangeLower = rangeUpper + 1 - rangeLower
}
}
// compute content length
var contentLength = uint64(fileInfo.Size)
if isRangeRead {
contentLength = rangeUpper - rangeLower + 1
w.WriteHeader(http.StatusPartialContent)
}
// get object tagging size