-
Notifications
You must be signed in to change notification settings - Fork 0
Mitlib extend custom request headers #1
New issue
Have a question about this project? Sign up for a free GitHub account to open an issue and contact its maintainers and the community.
By clicking “Sign up for GitHub”, you agree to our terms of service and privacy statement. We’ll occasionally send you account related emails.
Already on GitHub? Sign in to your account
base: main
Are you sure you want to change the base?
Changes from all commits
File filter
Filter by extension
Conversations
Jump to
Diff view
Diff view
There are no files selected for viewing
| Original file line number | Diff line number | Diff line change |
|---|---|---|
|
|
@@ -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,32 +391,43 @@ 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) | ||
| return r | ||
|
|
||
| @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) | ||
|
There was a problem hiding this comment. Choose a reason for hiding this commentThe reason will be displayed to describe this comment to others. Learn more. Appreciate the addition of named args here!
Author
There was a problem hiding this comment. Choose a reason for hiding this commentThe reason will be displayed to describe this comment to others. Learn more. Thanks, I thought so too! I much prefer named over positional when it gets even remotely complex, or the arg values look similar. |
||
| 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"] | ||
|
|
||
There was a problem hiding this comment.
Choose a reason for hiding this comment
The reason will be displayed to describe this comment to others. Learn more.
Great context!