diff --git a/tableauserverclient/models/view_item.py b/tableauserverclient/models/view_item.py index 146f21077..01635349b 100644 --- a/tableauserverclient/models/view_item.py +++ b/tableauserverclient/models/view_item.py @@ -1,5 +1,5 @@ import copy -from typing import Callable, Iterable, List, Optional, Set, TYPE_CHECKING +from typing import Callable, Generator, Iterator, List, Optional, Set, TYPE_CHECKING from defusedxml.ElementTree import fromstring @@ -24,8 +24,8 @@ def __init__(self) -> None: self._preview_image: Optional[Callable[[], bytes]] = None self._project_id: Optional[str] = None self._pdf: Optional[Callable[[], bytes]] = None - self._csv: Optional[Callable[[], Iterable[bytes]]] = None - self._excel: Optional[Callable[[], Iterable[bytes]]] = None + self._csv: Optional[Callable[[], Iterator[bytes]]] = None + self._excel: Optional[Callable[[], Iterator[bytes]]] = None self._total_views: Optional[int] = None self._sheet_type: Optional[str] = None self._updated_at: Optional["datetime"] = None @@ -94,14 +94,14 @@ def pdf(self) -> bytes: return self._pdf() @property - def csv(self) -> Iterable[bytes]: + def csv(self) -> Iterator[bytes]: if self._csv is None: error = "View item must be populated with its csv first." raise UnpopulatedPropertyError(error) return self._csv() @property - def excel(self) -> Iterable[bytes]: + def excel(self) -> Iterator[bytes]: if self._excel is None: error = "View item must be populated with its excel first." raise UnpopulatedPropertyError(error) diff --git a/tableauserverclient/server/endpoint/datasources_endpoint.py b/tableauserverclient/server/endpoint/datasources_endpoint.py index cb5600938..022523aa4 100644 --- a/tableauserverclient/server/endpoint/datasources_endpoint.py +++ b/tableauserverclient/server/endpoint/datasources_endpoint.py @@ -339,7 +339,7 @@ def update_hyper_data( *, request_id: str, actions: Sequence[Mapping], - payload: Optional[FilePath] = None + payload: Optional[FilePath] = None, ) -> JobItem: if isinstance(datasource_or_connection_item, DatasourceItem): datasource_id = datasource_or_connection_item.id diff --git a/tableauserverclient/server/endpoint/endpoint.py b/tableauserverclient/server/endpoint/endpoint.py index 8fdb74751..4260475ee 100644 --- a/tableauserverclient/server/endpoint/endpoint.py +++ b/tableauserverclient/server/endpoint/endpoint.py @@ -2,6 +2,7 @@ from distutils.version import LooseVersion as Version from functools import wraps from xml.etree.ElementTree import ParseError +from typing import Any, Callable, Dict, Optional, TYPE_CHECKING from .exceptions import ( ServerResponseError, @@ -18,9 +19,13 @@ XML_CONTENT_TYPE = "text/xml" JSON_CONTENT_TYPE = "application/json" +if TYPE_CHECKING: + from ..server import Server + from requests import Response + class Endpoint(object): - def __init__(self, parent_srv): + def __init__(self, parent_srv: "Server"): self.parent_srv = parent_srv @staticmethod @@ -46,13 +51,13 @@ def _safe_to_log(server_response): def _make_request( self, - method, - url, - content=None, - auth_token=None, - content_type=None, - parameters=None, - ): + method: Callable[..., "Response"], + url: str, + content: Optional[bytes] = None, + auth_token: Optional[str] = None, + content_type: Optional[str] = None, + parameters: Optional[Dict[str, Any]] = None, + ) -> "Response": parameters = parameters or {} parameters.update(self.parent_srv.http_options) if not "headers" in parameters: @@ -64,12 +69,23 @@ def _make_request( logger.debug("request {}, url: {}".format(method.__name__, url)) if content: - logger.debug("request content: {}".format(content[:1000])) + logger.debug("request content: %r", content[:1000]) server_response = method(url, **parameters) - self.parent_srv._namespace.detect(server_response.content) + + # Check if response is xml, and if so, validate namespace. Checking the + # content type header prevents eager evaluation of streaming requests. + if server_response.headers.get("Content-Type") == "application/xml": + self.parent_srv._namespace.detect(server_response.content) self._check_status(server_response) + # Response.content is a property. Calling it will load the entire response into memory. Checking if the + # content-type is an octet-stream accomplishes the same goal without eagerly loading content. + stream = parameters.get("stream", False) + stream = stream or (server_response.headers.get("Content-Type") == "application/octet-stream") + if stream: + return server_response + # This check is to determine if the response is a text response (xml or otherwise) # so that we do not attempt to log bytes and other binary data. if len(server_response.content) > 0 and server_response.encoding: diff --git a/tableauserverclient/server/endpoint/views_endpoint.py b/tableauserverclient/server/endpoint/views_endpoint.py index cb652fbc0..67e66a81f 100644 --- a/tableauserverclient/server/endpoint/views_endpoint.py +++ b/tableauserverclient/server/endpoint/views_endpoint.py @@ -9,7 +9,7 @@ logger = logging.getLogger("tableau.endpoint.views") -from typing import Iterable, List, Optional, Tuple, TYPE_CHECKING +from typing import Iterator, List, Optional, Tuple, TYPE_CHECKING if TYPE_CHECKING: from ..request_options import RequestOptions, CSVRequestOptions, PDFRequestOptions, ImageRequestOptions @@ -119,12 +119,11 @@ def csv_fetcher(): view_item._set_csv(csv_fetcher) logger.info("Populated csv for view (ID: {0})".format(view_item.id)) - def _get_view_csv(self, view_item: ViewItem, req_options: Optional["CSVRequestOptions"]) -> Iterable[bytes]: + def _get_view_csv(self, view_item: ViewItem, req_options: Optional["CSVRequestOptions"]) -> Iterator[bytes]: url = "{0}/{1}/data".format(self.baseurl, view_item.id) with closing(self.get_request(url, request_object=req_options, parameters={"stream": True})) as server_response: - csv = server_response.iter_content(1024) - return csv + yield from server_response.iter_content(1024) @api(version="3.8") def populate_excel(self, view_item: ViewItem, req_options: Optional["CSVRequestOptions"] = None) -> None: @@ -138,12 +137,11 @@ def excel_fetcher(): view_item._set_excel(excel_fetcher) logger.info("Populated excel for view (ID: {0})".format(view_item.id)) - def _get_view_excel(self, view_item: ViewItem, req_options: Optional["CSVRequestOptions"]) -> Iterable[bytes]: + def _get_view_excel(self, view_item: ViewItem, req_options: Optional["CSVRequestOptions"]) -> Iterator[bytes]: url = "{0}/{1}/crosstab/excel".format(self.baseurl, view_item.id) with closing(self.get_request(url, request_object=req_options, parameters={"stream": True})) as server_response: - excel = server_response.iter_content(1024) - return excel + yield from server_response.iter_content(1024) @api(version="3.2") def populate_permissions(self, item: ViewItem) -> None: diff --git a/test/test_endpoint.py b/test/test_endpoint.py new file mode 100644 index 000000000..d5321583b --- /dev/null +++ b/test/test_endpoint.py @@ -0,0 +1,28 @@ +from pathlib import Path +import unittest + +import tableauserverclient as TSC + +import requests_mock + +ASSETS = Path(__file__).parent / "assets" + + +class TestEndpoint(unittest.TestCase): + def setUp(self) -> None: + self.server = TSC.Server("http://test/", use_server_version=False) + + # Fake signin + self.server._site_id = "dad65087-b08b-4603-af4e-2887b8aafc67" + self.server._auth_token = "j80k54ll2lfMZ0tv97mlPvvSCRyD0DOM" + + return super().setUp() + + def test_get_request_stream(self) -> None: + url = "http://test/" + endpoint = TSC.server.Endpoint(self.server) + with requests_mock.mock() as m: + m.get(url, headers={"Content-Type": "application/octet-stream"}) + response = endpoint.get_request(url, parameters={"stream": True}) + + self.assertFalse(response._content_consumed)