Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
10 changes: 10 additions & 0 deletions README.md
Original file line number Diff line number Diff line change
Expand Up @@ -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"})
```

Comment on lines +49 to +58

Copy link
Copy Markdown

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Great context!

### Usage examples

See the `example.py` script for an example of community, collection, item, bundle and bitstream creation.
Expand Down
101 changes: 39 additions & 62 deletions dspace_rest_client/client.py
Original file line number Diff line number Diff line change
Expand Up @@ -24,7 +24,6 @@
from uuid import UUID

import requests
from requests import Request
import pysolr
import smart_open
from typing import cast, IO
Expand Down Expand Up @@ -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
Expand All @@ -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()
Expand All @@ -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):
Expand Down Expand Up @@ -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):
Expand All @@ -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
Expand Down Expand Up @@ -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)

Copy link
Copy Markdown

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Appreciate the addition of named args here!

Copy link
Copy Markdown
Author

Choose a reason for hiding this comment

The 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)
Expand Down Expand Up @@ -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
Expand All @@ -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.
Expand All @@ -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))
Expand All @@ -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

Expand Down Expand Up @@ -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"]
Expand Down
Loading