diff --git a/README.md b/README.md
index d2dc8a91..7494f0cd 100644
--- a/README.md
+++ b/README.md
@@ -274,7 +274,7 @@ curl -H "x-retry-test-id: 1d05c20627844214a9ff7cbcf696317d" "http://localhost:91
| return-broken-stream | [HTTP] Testbench will fail after a few downloaded bytes
[GRPC] Testbench will fail with `UNAVAILABLE` after a few downloaded bytes
| return-broken-stream-after-YK | [HTTP] Testbench will fail after YKiB of downloaded data
[GRPC] Testbench will fail with `UNAVAILABLE` after YKiB of downloaded data
| return-reset-connection | [HTTP] Testbench will fail with a reset connection
[GRPC] Testbench will fail the RPC with `UNAVAILABLE`
-| stall-for-Ts-after-YK | [HTTP] Testbench will stall for T second after reading YKiB of downloaded/uploaded data, e.g. stall-for-10s-after-12K stalls after reading/writing 12KiB of data
[GRPC] Not supported
+| stall-for-Ts-after-YK | [HTTP] Testbench will stall for T second after reading YKiB of downloaded/uploaded data, e.g. stall-for-10s-after-12K stalls after reading/writing 12KiB of data
[GRPC] Supported for `storage.objects.get` and `storage.objects.insert`
| redirect-send-token-T | [HTTP] Unsupported [GRPC] Testbench will fail the RPC with `ABORTED` and include appropriate redirection error details.
| redirect-send-handle-and-token-T | [HTTP] Unsupported [GRPC] Testbench will fail the RPC with `ABORTED` and include appropriate redirection error details.
| return-X-if-dp-enforced | [HTTP] Unsupported [GRPC] Testbench will fail with the equivalent gRPC error to the HTTP code provided for X, but only if DirectPath is enforced.
diff --git a/gcs/upload.py b/gcs/upload.py
index f7a44541..573b4497 100644
--- a/gcs/upload.py
+++ b/gcs/upload.py
@@ -269,6 +269,24 @@ def init_write_object_grpc(cls, db, request_iterator, context):
test_id=test_id,
)
+ # Handle retry test stall-for-Xs-after-YK instructions if applicable.
+ (
+ stall_time,
+ after_bytes,
+ test_id,
+ ) = testbench.common.get_stall_uploads_after_bytes(
+ db, request, context=context, transport="GRPC"
+ )
+ if stall_time:
+ testbench.common.handle_stall_uploads_after_bytes(
+ upload,
+ content,
+ db,
+ stall_time,
+ after_bytes,
+ test_id=test_id,
+ )
+
upload.media += content
if request.finish_write:
upload.complete = True
@@ -582,6 +600,24 @@ def update_upload_checksums(upload_metadata, object_checksums):
test_id=test_id,
)
+ # Handle retry test stall-for-Xs-after-YK instructions if applicable.
+ (
+ stall_time,
+ after_bytes,
+ test_id,
+ ) = testbench.common.get_stall_uploads_after_bytes(
+ db, request, context=context, transport="GRPC"
+ )
+ if stall_time:
+ testbench.common.handle_stall_uploads_after_bytes(
+ upload,
+ content,
+ db,
+ stall_time,
+ after_bytes,
+ test_id=test_id,
+ )
+
# Currently, the testbench will always checkpoint and flush data for testing purposes,
# instead of the 15 seconds interval used in the GCS server.
# TODO(#592): Refactor testbench checkpointing to more closely follow GCS server behavior.
diff --git a/testbench/common.py b/testbench/common.py
index 97378ddd..56b62feb 100644
--- a/testbench/common.py
+++ b/testbench/common.py
@@ -929,7 +929,10 @@ def wrapper(*args, **kwargs):
def get_stall_uploads_after_bytes(database, request, context=None, transport="HTTP"):
"""Retrieve stall time and #bytes corresponding to uploads from retry test instructions."""
method = "storage.objects.insert"
- test_id = request.headers.get("x-retry-test-id", None)
+ if context is not None:
+ test_id = get_retry_test_id_from_context(context)
+ else:
+ test_id = request.headers.get("x-retry-test-id", None)
if not test_id:
return 0, 0, ""
next_instruction = None
@@ -992,9 +995,14 @@ def handle_stall_uploads_after_bytes(
e.g. We are uploading 120K of data then, stall-2s-after-100K will stall the request.
"""
if len(upload.media) <= after_bytes and len(upload.media) + len(data) > after_bytes:
+ should_stall = True
if test_id:
- database.dequeue_next_instruction(test_id, "storage.objects.insert")
- time.sleep(stall_time)
+ should_stall = (
+ database.dequeue_next_instruction(test_id, "storage.objects.insert")
+ is not None
+ )
+ if should_stall:
+ time.sleep(stall_time)
def handle_retry_uploads_error_after_bytes(
diff --git a/testbench/database.py b/testbench/database.py
index f47f9081..d9653f48 100644
--- a/testbench/database.py
+++ b/testbench/database.py
@@ -752,7 +752,9 @@ def insert_retry_test(self, instructions, transport="HTTP"):
def has_instructions_retry_test(self, retry_test_id, method, transport="HTTP"):
with self._retry_tests_lock:
- retry_test = self.get_retry_test(retry_test_id)
+ retry_test = self._retry_tests.get(retry_test_id, None)
+ if retry_test is None:
+ return False
# Add validation for request transport as well.
if (len(retry_test["instructions"].get(method, [])) > 0) and retry_test[
"transport"
diff --git a/testbench/grpc_server.py b/testbench/grpc_server.py
index 4f8bbdcf..533810ad 100644
--- a/testbench/grpc_server.py
+++ b/testbench/grpc_server.py
@@ -20,6 +20,7 @@
import json
import re
import sys
+import time
import types
import uuid
from collections.abc import Iterable
@@ -42,7 +43,38 @@
from google.storage.control.v2 import storage_control_pb2, storage_control_pb2_grpc
from google.storage.v2 import storage_pb2, storage_pb2_grpc
-_GRPC_SERVER_THREAD_COUNT = 2
+_GRPC_SERVER_THREAD_COUNT = 8
+
+
+def _should_stall_after_bytes(bytes_yielded, chunk_len, stall_after_bytes):
+ if chunk_len <= 0:
+ return False
+ if stall_after_bytes == 0:
+ return bytes_yielded == 0
+ return bytes_yielded < stall_after_bytes <= bytes_yielded + chunk_len
+
+
+def _apply_grpc_read_stall_if_applicable(
+ database,
+ test_id,
+ method,
+ bytes_yielded,
+ chunk_len,
+ stall_time,
+ stall_after_bytes,
+):
+ if not test_id:
+ return False
+
+ if stall_time <= 0 or not _should_stall_after_bytes(
+ bytes_yielded, chunk_len, stall_after_bytes
+ ):
+ return False
+
+ if database.dequeue_next_instruction(test_id, method) is None:
+ return False
+
+ time.sleep(stall_time)
def _trimmed_content(content):
@@ -607,6 +639,9 @@ def ReadObject(self, request, context):
# Check retry test broken-stream instructions.
test_id = testbench.common.get_retry_test_id_from_context(context)
broken_stream_after_bytes = 0
+ stall_time = 0
+ stall_after_bytes = 0
+ stall_applied = False
method = "storage.objects.get"
if test_id and self.db.has_instructions_retry_test(
test_id, method, transport="GRPC"
@@ -615,12 +650,33 @@ def ReadObject(self, request, context):
broken_stream_after_bytes = testbench.common.get_broken_stream_after_bytes(
next_instruction
)
+ retry_stall_after_bytes_matches = (
+ testbench.common.retry_stall_after_bytes.match(next_instruction)
+ )
+ if retry_stall_after_bytes_matches:
+ items = list(retry_stall_after_bytes_matches.groups())
+ stall_time = int(items[0])
+ stall_after_bytes = int(items[1]) * 1024
+ bytes_yielded = 0
while start <= read_end:
end = min(start + size, read_end)
+ chunk_len = end - start
+
+ _apply_grpc_read_stall_if_applicable(
+ self.db,
+ test_id,
+ method,
+ bytes_yielded,
+ chunk_len,
+ stall_time,
+ stall_after_bytes,
+ )
+
# Handle retry test broken-stream failures if applicable.
if broken_stream_after_bytes and end >= broken_stream_after_bytes:
chunk = blob.media[start:broken_stream_after_bytes]
+ bytes_yielded += len(chunk)
yield storage_pb2.ReadObjectResponse(
checksummed_data={
"content": chunk,
@@ -636,6 +692,7 @@ def ReadObject(self, request, context):
"Injected 'broken stream' fault",
)
chunk = blob.media[start:end]
+ bytes_yielded += len(chunk)
yield storage_pb2.ReadObjectResponse(
checksummed_data={
"content": chunk,
@@ -670,6 +727,9 @@ def BidiReadObject(self, request_iterator, context):
# Check retry test broken-stream instructions.
test_id = testbench.common.get_retry_test_id_from_context(context)
broken_stream_after_bytes = 0
+ stall_time = 0
+ stall_after_bytes = 0
+ stall_applied = False
method = "storage.objects.get"
if test_id and self.db.has_instructions_retry_test(
test_id, method, transport="GRPC"
@@ -678,6 +738,13 @@ def BidiReadObject(self, request_iterator, context):
broken_stream_after_bytes = testbench.common.get_broken_stream_after_bytes(
next_instruction
)
+ retry_stall_after_bytes_matches = (
+ testbench.common.retry_stall_after_bytes.match(next_instruction)
+ )
+ if retry_stall_after_bytes_matches:
+ items = list(retry_stall_after_bytes_matches.groups())
+ stall_time = int(items[0])
+ stall_after_bytes = int(items[1]) * 1024
return_redirect_token = (
testbench.common.get_return_read_handle_and_redirect_token(self.db, context)
)
@@ -782,6 +849,7 @@ def read_results():
for request in request_iterator:
yield from responses_for_range_batch(request.read_ranges)
+ bytes_yielded = 0
for chunk, range_end, read_range in read_results():
count = len(chunk)
excess = count - returnable
@@ -789,6 +857,18 @@ def read_results():
chunk = chunk[:returnable]
range_end = False
read_range["read_length"] -= excess
+
+ _apply_grpc_read_stall_if_applicable(
+ self.db,
+ test_id,
+ method,
+ bytes_yielded,
+ len(chunk),
+ stall_time,
+ stall_after_bytes,
+ )
+
+ bytes_yielded += len(chunk)
returnable -= count
yield response(
storage_pb2.BidiReadObjectResponse(
diff --git a/testbench/rest_server.py b/testbench/rest_server.py
index ca6c123d..b08515fe 100644
--- a/testbench/rest_server.py
+++ b/testbench/rest_server.py
@@ -1247,10 +1247,10 @@ def delete_resumable_upload(bucket_name):
# === SERVER === #
# Define the WSGI application to handle HMAC key and service account requests
-(PROJECTS_HANDLER_PATH, projects_app) = projects_rest_server.get_projects_app(db)
+PROJECTS_HANDLER_PATH, projects_app = projects_rest_server.get_projects_app(db)
# Define the WSGI application to handle IAM requests
-(IAM_HANDLER_PATH, iam_app) = iam_rest_server.get_iam_app()
+IAM_HANDLER_PATH, iam_app = iam_rest_server.get_iam_app()
server = flask.Flask(__name__)
server.debug = False
diff --git a/tests/test_testbench_retry.py b/tests/test_testbench_retry.py
index 6ea3c19e..4d7d6959 100644
--- a/tests/test_testbench_retry.py
+++ b/tests/test_testbench_retry.py
@@ -1062,6 +1062,206 @@ def test_grpc_retry_reset_connection(self):
"Injected 'socket closed, connection reset by peer' fault",
)
+ def test_grpc_retry_stall_read_after_bytes(self):
+ media = self._create_block(2 * UPLOAD_QUANTUM)
+ response = self.rest_client.put(
+ "/bucket-name/512k.txt",
+ content_type="text/plain",
+ data=media,
+ )
+ self.assertEqual(response.status_code, 200)
+
+ response = self.rest_client.post(
+ "/retry_test",
+ data=json.dumps(
+ {
+ "instructions": {
+ "storage.objects.get": ["stall-for-1s-after-128K"]
+ },
+ "transport": "GRPC",
+ },
+ ),
+ )
+ self.assertEqual(response.status_code, 200)
+ create_rest = json.loads(response.data)
+ self.assertIn("id", create_rest)
+
+ context = unittest.mock.Mock()
+ context.invocation_metadata = unittest.mock.Mock(
+ return_value=(("x-retry-test-id", create_rest.get("id")),)
+ )
+
+ start_time = time.perf_counter()
+ response = self.grpc.ReadObject(
+ storage_pb2.ReadObjectRequest(
+ bucket="projects/_/buckets/bucket-name", object="512k.txt"
+ ),
+ context,
+ )
+ list(response)
+ elapsed = time.perf_counter() - start_time
+ self.assertGreater(elapsed, 1)
+
+ def test_grpc_retry_stall_write_after_bytes(self):
+ response = self.rest_client.post(
+ "/retry_test",
+ data=json.dumps(
+ {
+ "instructions": {
+ "storage.objects.insert": ["stall-for-1s-after-250K"]
+ },
+ "transport": "GRPC",
+ }
+ ),
+ )
+ self.assertEqual(response.status_code, 200)
+ create_rest = json.loads(response.data)
+ self.assertIn("id", create_rest)
+ id = create_rest.get("id")
+
+ context = unittest.mock.Mock()
+ context.invocation_metadata = unittest.mock.Mock(
+ return_value=(("x-retry-test-id", id),)
+ )
+ start = self.grpc.StartResumableWrite(
+ storage_pb2.StartResumableWriteRequest(
+ write_object_spec=storage_pb2.WriteObjectSpec(
+ resource=storage_pb2.Object(
+ name="object-name-stall",
+ bucket="projects/_/buckets/bucket-name",
+ )
+ )
+ ),
+ context=context,
+ )
+ self.assertIsNotNone(start.upload_id)
+
+ content = self._create_block(UPLOAD_QUANTUM).encode("utf-8")
+ r1 = storage_pb2.WriteObjectRequest(
+ upload_id=start.upload_id,
+ write_offset=0,
+ checksummed_data=storage_pb2.ChecksummedData(
+ content=content, crc32c=crc32c.crc32c(content)
+ ),
+ finish_write=False,
+ )
+ start_time = time.perf_counter()
+ _ = self.grpc.WriteObject([r1], context)
+ elapsed = time.perf_counter() - start_time
+ self.assertGreater(elapsed, 1)
+
+ # Instruction consumed; finishing write should be fast.
+ r2 = storage_pb2.WriteObjectRequest(
+ upload_id=start.upload_id,
+ write_offset=len(content),
+ checksummed_data=storage_pb2.ChecksummedData(
+ content=b"", crc32c=crc32c.crc32c(b"")
+ ),
+ finish_write=True,
+ )
+ start_time = time.perf_counter()
+ _ = self.grpc.WriteObject([r2], context)
+ elapsed = time.perf_counter() - start_time
+ self.assertLess(elapsed, 1)
+
+ def test_grpc_retry_stall_bidiwrite_after_bytes(self):
+ response = self.rest_client.post(
+ "/retry_test",
+ data=json.dumps(
+ {
+ "instructions": {
+ "storage.objects.insert": ["stall-for-1s-after-250K"]
+ },
+ "transport": "GRPC",
+ }
+ ),
+ )
+ self.assertEqual(response.status_code, 200)
+ create_rest = json.loads(response.data)
+ self.assertIn("id", create_rest)
+ id = create_rest.get("id")
+
+ context = unittest.mock.Mock()
+ context.invocation_metadata = unittest.mock.Mock(
+ return_value=(("x-retry-test-id", id),)
+ )
+ start = self.grpc.StartResumableWrite(
+ storage_pb2.StartResumableWriteRequest(
+ write_object_spec=storage_pb2.WriteObjectSpec(
+ resource=storage_pb2.Object(
+ name="object-name-bidi-stall",
+ bucket="projects/_/buckets/bucket-name",
+ )
+ )
+ ),
+ context=context,
+ )
+ self.assertIsNotNone(start.upload_id)
+
+ content = self._create_block(UPLOAD_QUANTUM).encode("utf-8")
+ r1 = storage_pb2.BidiWriteObjectRequest(
+ upload_id=start.upload_id,
+ write_offset=0,
+ checksummed_data=storage_pb2.ChecksummedData(
+ content=content, crc32c=crc32c.crc32c(content)
+ ),
+ finish_write=False,
+ )
+
+ start_time = time.perf_counter()
+ _ = list(self.grpc.BidiWriteObject([r1], context))
+ elapsed = time.perf_counter() - start_time
+ self.assertGreater(elapsed, 1)
+
+ def test_grpc_bidiread_retry_stall_after_bytes(self):
+ media = self._create_block(5 * 1024 * 1024)
+ response = self.rest_client.put(
+ "/bucket-name/512k.txt",
+ content_type="text/plain",
+ data=media,
+ )
+ self.assertEqual(response.status_code, 200)
+
+ response = self.rest_client.post(
+ "/retry_test",
+ data=json.dumps(
+ {
+ "instructions": {
+ "storage.objects.get": ["stall-for-1s-after-256K"]
+ },
+ "transport": "GRPC",
+ },
+ ),
+ )
+ self.assertEqual(response.status_code, 200)
+ create_rest = json.loads(response.data)
+ self.assertIn("id", create_rest)
+
+ context = unittest.mock.Mock()
+ context.invocation_metadata = unittest.mock.Mock(
+ return_value=(("x-retry-test-id", create_rest.get("id")),)
+ )
+
+ r1 = storage_pb2.BidiReadObjectRequest(
+ read_object_spec=storage_pb2.BidiReadObjectSpec(
+ bucket="projects/_/buckets/bucket-name",
+ object="512k.txt",
+ ),
+ read_ranges=[
+ storage_pb2.ReadRange(
+ read_offset=0,
+ read_length=1 * 1024 * 1024,
+ read_id=1,
+ ),
+ ],
+ )
+
+ start_time = time.perf_counter()
+ response = self.grpc.BidiReadObject([r1], context)
+ list(response)
+ elapsed = time.perf_counter() - start_time
+ self.assertGreater(elapsed, 1)
+
def test_grpc_retry_broken_stream(self):
# Use the XML API to inject an object with some data.
media = self._create_block(2 * UPLOAD_QUANTUM)