diff --git a/tableauserverapi/models/tableau_auth.py b/tableauserverapi/models/tableau_auth.py index f6b98fc38..cdf0fb410 100644 --- a/tableauserverapi/models/tableau_auth.py +++ b/tableauserverapi/models/tableau_auth.py @@ -1,19 +1,6 @@ -import xml.etree.ElementTree as ET -from .. import NAMESPACE - - class TableauAuth(object): - def __init__(self, username, password, site='', impersonate_id=None): - # CHECK FOR USERNAME AND PASSWORD + def __init__(self, username, password, site='', user_id_to_impersonate=None): + self.user_id_to_impersonate = user_id_to_impersonate self.password = password - self.username = username self.site = site - self.impersonate_id = impersonate_id - - @staticmethod - def from_response(parent_srv, resp): - parsed_response = ET.fromstring(resp) - parent_srv._site_id = parsed_response.find('.//t:site', namespaces=NAMESPACE).get('id', None) - parent_srv._user_id = parsed_response.find('.//t:user', namespaces=NAMESPACE).get('id', None) - auth_token = parsed_response.find('t:credentials', namespaces=NAMESPACE).get('token', None) - parent_srv._auth_token = auth_token + self.username = username diff --git a/tableauserverapi/server/endpoint/auth_endpoint.py b/tableauserverapi/server/endpoint/auth_endpoint.py index 55517c26d..816ba7ceb 100644 --- a/tableauserverapi/server/endpoint/auth_endpoint.py +++ b/tableauserverapi/server/endpoint/auth_endpoint.py @@ -1,5 +1,6 @@ from endpoint import Endpoint -from .. import RequestFactory, TableauAuth +from .. import RequestFactory, NAMESPACE +import xml.etree.ElementTree as ET import logging logger = logging.getLogger('tableau.endpoint.auth') @@ -27,7 +28,11 @@ def sign_in(self, auth_req): server_response = self.parent_srv.session.post(url, data=signin_req, **self.parent_srv.http_options) Endpoint._check_status(server_response) - TableauAuth.from_response(self.parent_srv, server_response.text) + parsed_response = ET.fromstring(server_response.text) + site_id = parsed_response.find('.//t:site', namespaces=NAMESPACE).get('id', None) + user_id = parsed_response.find('.//t:user', namespaces=NAMESPACE).get('id', None) + auth_token = parsed_response.find('t:credentials', namespaces=NAMESPACE).get('token', None) + self.parent_srv._set_auth(site_id, user_id, auth_token) logger.info('Signed into {0} as {1}'.format(self.parent_srv.server_address, auth_req.username)) return Auth.contextmgr(self.sign_out) diff --git a/tableauserverapi/server/request_factory.py b/tableauserverapi/server/request_factory.py index 2acd4db84..e39d9ac81 100644 --- a/tableauserverapi/server/request_factory.py +++ b/tableauserverapi/server/request_factory.py @@ -10,9 +10,9 @@ def _add_multipart(parts): multipart_part = RequestField(name=name, data=data, filename=filename) multipart_part.make_multipart(content_type=content_type) mime_multipart_parts.append(multipart_part) - post_body, content_type = encode_multipart_formdata(mime_multipart_parts) + xml_request, content_type = encode_multipart_formdata(mime_multipart_parts) content_type = ''.join(('multipart/mixed',) + content_type.partition(';')[1:]) - return post_body, content_type + return xml_request, content_type class AuthRequest(object): @@ -23,9 +23,9 @@ def signin_req(self, auth_item): credentials_element.attrib['password'] = auth_item.password site_element = ET.SubElement(credentials_element, 'site') site_element.attrib['contentUrl'] = auth_item.site - if auth_item.impersonate_id: + if auth_item.user_id_to_impersonate: user_element = ET.SubElement(credentials_element, 'user') - user_element.attrib['id'] = auth_item.impersonate_id + user_element.attrib['id'] = auth_item.user_id_to_impersonate return ET.tostring(xml_request) diff --git a/tableauserverapi/server/server.py b/tableauserverapi/server/server.py index fd03f2d02..9c8b0212f 100644 --- a/tableauserverapi/server/server.py +++ b/tableauserverapi/server/server.py @@ -34,11 +34,16 @@ def clear_http_options(self): self._http_options = dict() def _clear_auth(self): - self._auth_token = None self._site_id = None self._user_id = None + self._auth_token = None self._session = requests.Session() + def _set_auth(self, site_id, user_id, auth_token): + self._site_id = site_id + self._user_id = user_id + self._auth_token = auth_token + @property def baseurl(self): return "{0}/api/{1}".format(self._server_address, str(self.version)) diff --git a/test/test_auth.py b/test/test_auth.py index 2efcb688d..36482b21a 100644 --- a/test/test_auth.py +++ b/test/test_auth.py @@ -32,9 +32,8 @@ def test_sign_in_impersonate(self): response_xml = file.read() with requests_mock.mock() as m: m.post(self.baseurl + '/signin', text=response_xml) - tableau_auth = TSA.TableauAuth('testuser', - 'password', - impersonate_id='dd2239f6-ddf1-4107-981a-4cf94e415794') + tableau_auth = TSA.TableauAuth('testuser', 'password', + user_id_to_impersonate='dd2239f6-ddf1-4107-981a-4cf94e415794') self.server.auth.sign_in(tableau_auth) self.assertEqual('MJonFA6HDyy2C3oqR13fRGqE6cmgzwq3', self.server.auth_token)