diff --git a/azure/functions/_http_wsgi.py b/azure/functions/_http_wsgi.py index 58c3c0e..9ed5d34 100644 --- a/azure/functions/_http_wsgi.py +++ b/azure/functions/_http_wsgi.py @@ -151,7 +151,12 @@ def __init__(self): @classmethod def from_app(cls, app, environ) -> 'WsgiResponse': res = cls() - res._buffer = [x or b'' for x in app(environ, res._start_response)] + response = app(environ, res._start_response) + try: + res._buffer = [x or b'' for x in response] + finally: + if hasattr(response, 'close'): + response.close() return res def to_func_response(self) -> HttpResponse: diff --git a/tests/test_http_wsgi.py b/tests/test_http_wsgi.py index d3fecec..1d465c6 100644 --- a/tests/test_http_wsgi.py +++ b/tests/test_http_wsgi.py @@ -5,6 +5,7 @@ from io import StringIO, BytesIO import pytest +from werkzeug.wsgi import ClosingIterator, FileWrapper import azure.functions as func from azure.functions._abc import TraceContext, RetryContext @@ -24,6 +25,45 @@ def __init__(self, message=''): class TestHttpWsgi(unittest.TestCase): + def test_response_closes_file_wrapper(self): + stream = BytesIO(b'file contents') + + def app(environ, start_response): + start_response('200 OK', []) + return FileWrapper(stream, buffer_size=4) + + response = WsgiResponse.from_app(app, {}).to_func_response() + self.assertEqual(response.get_body(), b'file contents') + self.assertTrue(stream.closed) + + def test_response_closes_empty_iterable(self): + closed = [] + + def app(environ, start_response): + start_response('204 No Content', []) + return ClosingIterator([], lambda: closed.append(True)) + + response = WsgiResponse.from_app(app, {}).to_func_response() + self.assertEqual(response.get_body(), b'') + self.assertEqual(closed, [True]) + + def test_response_closes_iterable_after_iteration_error(self): + closed = [] + error = WsgiException('stream failed') + + def chunks(): + yield b'partial response' + raise error + + def app(environ, start_response): + start_response('200 OK', []) + return ClosingIterator(chunks(), lambda: closed.append(True)) + + with self.assertRaises(WsgiException) as caught: + WsgiResponse.from_app(app, {}) + self.assertIs(caught.exception, error) + self.assertEqual(closed, [True]) + def test_request_general_environ_conversion(self): func_request = self._generate_func_request() error_buffer = StringIO()