diff --git a/README.md b/README.md index 8c171ab..bda8e5a 100644 --- a/README.md +++ b/README.md @@ -46,6 +46,16 @@ DSpace-CRIS allows API tokens for users to be created in their EPerson profile p If a token is found, it will be set as the Authorization Bearer header instead of requesting JWT bearer tokens. +### Custom request headers + +Some environments require an additional header on every request (e.g. for a proxy or gateway). Pass these as +`request_header_mixins` when creating the client, and they will be included in all REST API requests, taking +precedence over the default headers: + +```python +d = DSpaceClient(request_header_mixins={"X-Foo": "bar"}) +``` + ### Usage examples See the `example.py` script for an example of community, collection, item, bundle and bitstream creation. diff --git a/dspace_rest_client/client.py b/dspace_rest_client/client.py index 1942ee8..dde7460 100644 --- a/dspace_rest_client/client.py +++ b/dspace_rest_client/client.py @@ -24,7 +24,6 @@ from uuid import UUID import requests -from requests import Request import pysolr import smart_open from typing import cast, IO @@ -223,6 +222,7 @@ def __init__( solr_auth=SOLR_AUTH, fake_user_agent=False, proxies=PROXY_DICT, + request_header_mixins=None, ): """ Accept optional API endpoint, username, password arguments using the OS environment @@ -231,6 +231,9 @@ def __init__( :param username: username with appropriate privileges to perform operations on REST API :param password: password for the above username + :param request_header_mixins: optional dict of headers to include in every REST API + request, eg. {"X-Foo": "bar"}. These take precedence over + the default headers. """ self.session = requests.Session() self.api_token = read_personal_api_token_secret() @@ -253,15 +256,21 @@ def __init__( "Mozilla/5.0 (Windows NT 6.2; WOW64) AppleWebKit/537.36 (KHTML, like Gecko) " "Chrome/39.0.2171.95 Safari/537.36" ) - # Set headers based on this - self.auth_request_headers = {"User-Agent": self.USER_AGENT} + # Set headers based on this, with any mixins applied last so they take precedence + self.request_header_mixins = dict(request_header_mixins or {}) + self.auth_request_headers = { + "User-Agent": self.USER_AGENT, + **self.request_header_mixins, + } self.request_headers = { "Content-type": "application/json", "User-Agent": self.USER_AGENT, + **self.request_header_mixins, } self.list_request_headers = { "Content-type": "text/uri-list", "User-Agent": self.USER_AGENT, + **self.request_header_mixins, } def authenticate(self, retry=False): @@ -373,7 +382,7 @@ def refresh_token(self): If the DSPACE-XSRF-TOKEN appears, we need to update our local stored token and re-send our API request @return: None """ - r = self.api_post(self.LOGIN_URL, None, None) + r = self.api_post(self.LOGIN_URL) self.update_token(r) def api_get(self, url, params=None, data=None, headers=None): @@ -382,13 +391,13 @@ def api_get(self, url, params=None, data=None, headers=None): @param url: DSpace REST API URL @param params: any parameters to include (eg ?page=0) @param data: any data to supply (typically not relevant for GET) - @param headers: any override headers (eg. with short-lived token for download) + @param headers: optional headers, merged over the default request headers + (eg. with short-lived token for download) @return: Response from API """ - if headers is None: - headers = self.request_headers + request_headers = {**self.request_headers, **(headers or {})} r = self.session.get(url, params=params, data=data, - headers=headers, + headers=request_headers, proxies=self.proxies ) self.update_token(r) @@ -396,18 +405,29 @@ def api_get(self, url, params=None, data=None, headers=None): @reauthenticate @refresh_csrf - def api_post(self, url, params, json): + def api_post( + self, url, *, params=None, json=None, data=None, files=None, headers=None + ): """ Perform a POST request. Refresh XSRF token if necessary. POSTs are typically used to create objects. @param url: DSpace REST API URL @param params: Any parameters to include (eg ?parent=abbc-....) @param json: Data in json-ready form (dict) to send as POST body (eg. item.as_dict()) + @param data: Form data to send as POST body (eg. multipart fields alongside files) + @param files: Files to send as a multipart upload, in the form accepted by requests + @param headers: Optional headers, merged over the default request headers @return: Response from API """ + request_headers = {**self.request_headers, **(headers or {})} + if files is not None: + # let requests set the multipart Content-Type, including the boundary + request_headers = { + k: v for k, v in request_headers.items() if k.lower() != "content-type" + } r = self.session.post( - url, json=json, params=params, headers=self.request_headers, - proxies=self.proxies + url, json=json, data=data, files=files, params=params, + headers=request_headers, proxies=self.proxies ) self.update_token(r) return r @@ -676,7 +696,7 @@ def create_dso(self, url, params, data, embeds=None): @return: Raw API response. New DSO *could* be returned but for error checking purposes, raw response is nice too and can always be parsed from this response later. """ - r = self.api_post(url, parse_params(params, embeds), data) + r = self.api_post(url, params=parse_params(params, embeds), json=data) if r.status_code == 201: # 201 Created - success! new_dso = parse_json(r) @@ -934,16 +954,12 @@ def create_bitstream( mime=None, metadata=None, embeds=None, - retry=False, - reauthenticated=False, ): """ Upload a file and create a bitstream for a specified parent bundle, from the uploaded file and the supplied metadata. - This create method is a bit different to the others, it does not use create_dso or the api_post lower level - methods, instead it has to use a prepared session POST request which will allow the multi-part upload to work - successfully with the correct byte size and persist the session data. - This is also why it directly implements the 'retry' functionality instead of relying on api_post. + The file and properties are sent as a multipart upload via api_post, which handles XSRF token + refreshes and reauthentication. @param bundle: python Bundle object @param name: Bitstream name @param path: Path to the file that will be uploaded. Can be a local filesystem path or any cloud @@ -952,10 +968,7 @@ def create_bitstream( storage authentication is handled outside of this application. @param mime: MIME string of the uploaded file @param metadata: Full metadata JSON - @param retry: A 'retried' indicator. If the first attempt fails due to an expired or missing auth - token, the request will retry once, after the token is refreshed. (default: False) - @param reauthenticated: Whether a reauthenticate attempt (in case of HTTP 401) was already attempted. - @return: constructed Bitstream object from the API response, or None if the operation failed. + @return: constructed Bitstream object from the API response, or None if the operation failed. """ # TODO: It is probably wise to allow the bundle UUID to be simply passed as an alternative to having the full # python object as constructed by this REST client, for more flexible usage. @@ -977,41 +990,9 @@ def create_bitstream( except Exception as e: logging.error("Error reading file from %s: %s", path, str(e)) return None - h = self.session.headers - h.update({"Content-Encoding": "gzip", "User-Agent": self.USER_AGENT}) - req = Request( - "POST", - url, - data=payload, - headers=h, - files=files, - params=parse_params(embeds=embeds), + r = self.api_post( + url, params=parse_params(embeds=embeds), data=payload, files=files ) - prepared_req = self.session.prepare_request(req) - r = self.session.send(prepared_req, proxies=self.proxies) - if "DSPACE-XSRF-TOKEN" in r.headers: - t = r.headers["DSPACE-XSRF-TOKEN"] - logging.debug("Updating token to %s", t) - self.session.headers.update({"X-XSRF-Token": t}) - self.session.cookies.update({"X-XSRF-Token": t}) - # as this method doesn't return the request, we cannot use our @refresh_csft decorator - # we should enhance self.api_post to be able to send files and use our decorators - if r.status_code == 403: - r_json = parse_json(r) - if r_json is not None and "message" in r_json and "CSRF token" in r_json["message"]: - if retry: - logging.error("Already retried... something must be wrong") - else: - logging.debug("Retrying request with updated CSRF token") - return self.create_bitstream( - bundle, name, path, mime, metadata, embeds, True - ) - # as this method doesn't return the request, we cannot use our @reauthenticate decorator - # we should enhance self.api_post to be able to send files and use our decorators - if r.status_code == 401 and not reauthenticated: - self.authenticate() - prepared_req = self.session.prepare_request(req) - r = self.session.send(prepared_req) if r.status_code == 201 or r.status_code == 200: # Success return Bitstream(api_resource=parse_json(r)) @@ -1026,11 +1007,7 @@ def download_bitstream(self, uuid=None): @return: full response object including headers, and content """ url = f"{self.API_ENDPOINT}/core/bitstreams/{uuid}/content" - h = { - "User-Agent": self.USER_AGENT, - "Authorization": self.get_short_lived_token(), - } - r = self.api_get(url, headers=h) + r = self.api_get(url, headers={"Authorization": self.get_short_lived_token()}) if r.status_code == 200: return r @@ -1564,7 +1541,7 @@ def get_short_lived_token(self): self.session = requests.Session() url = f"{self.API_ENDPOINT}/authn/shortlivedtokens" - r = self.api_post(url, json=None, params=None) + r = self.api_post(url) r_json = parse_json(r) if r_json is not None and "token" in r_json: return r_json["token"]