diff --git a/razorpay/resources/payment.py b/razorpay/resources/payment.py index ee5e3130..ad274fee 100644 --- a/razorpay/resources/payment.py +++ b/razorpay/resources/payment.py @@ -49,20 +49,37 @@ def capture(self, payment_id, amount, data={}, **kwargs): # nosemgrep : python.l data['amount'] = amount return self.post_url(url, data, **kwargs) - def refund(self, payment_id, amount, data={}, **kwargs): # pragma: no cover # nosemgrep : python.lang.correctness.common-mistakes.default-mutable-dict.default-mutable-dict + def refund(self, payment_id, data_or_amount=None, data=None, **kwargs): """ Refund Payment for given Id Args: payment_id : Id for which payment object has to be refunded - amount : Amount for which the payment has to be refunded + data_or_amount : Either the refund amount (int, float, str) or a dictionary + containing refund parameters (e.g. amount, notes, speed, receipt) + data : Optional dictionary of refund options when amount is passed as positional argument + **kwargs : Additional arguments including keyword 'amount' or request options Returns: Payment dict after getting refunded """ url = "{}/{}/refund".format(self.base_url, payment_id) - data['amount'] = amount - return self.post_url(url, data, **kwargs) + + payload = {} + if isinstance(data_or_amount, dict): + payload.update(data_or_amount) + elif data_or_amount is not None: + payload['amount'] = data_or_amount + + if isinstance(data, dict): + payload.update(data) + elif data is not None and 'amount' not in payload: + payload['amount'] = data + + if 'amount' in kwargs: + payload['amount'] = kwargs.pop('amount') + + return self.post_url(url, payload, **kwargs) def transfer(self, payment_id, data={}, **kwargs): """ @@ -116,16 +133,6 @@ def upi_transfer(self, payment_id, data={}, **kwargs): """ url = "{}/{}/upi_transfer".format(self.base_url, payment_id) return self.get_url(url, data, **kwargs) - - def refund(self, payment_id, data={}, **kwargs): - """ - Create a normal refund - - Returns: - Payment dict after getting refund - """ - url = "{}/{}/refund".format(self.base_url, payment_id) - return self.post_url(url, data, **kwargs) def fetch_multiple_refund(self, payment_id, data={}, **kwargs): """ @@ -139,7 +146,7 @@ def fetch_multiple_refund(self, payment_id, data={}, **kwargs): def fetch_refund_id(self, payment_id, refund_id, **kwargs): """ - Fetch multiple refunds for a payment + Fetch a specific refund by ID for a payment Returns: Refund dict diff --git a/tests/test_client_payment.py b/tests/test_client_payment.py index 0440f94f..ea94a2de 100644 --- a/tests/test_client_payment.py +++ b/tests/test_client_payment.py @@ -52,6 +52,8 @@ def test_refund_create(self): match_querystring=True) self.assertEqual(self.client.payment.refund(self.payment_id, 2000), result) + request_body = json.loads(responses.calls[0].request.body) + self.assertEqual(request_body, {'amount': 2000}) @responses.activate def test_transfer(self): @@ -116,7 +118,64 @@ def test_payment_refund(self): url = '{}/{}/refund'.format(self.base_url, 'fake_refund_id') responses.add(responses.POST, url, status=200, body=json.dumps(result), match_querystring=True) - self.assertEqual(self.client.payment.refund('fake_refund_id',init), result) + self.assertEqual(self.client.payment.refund('fake_refund_id',init), result) + request_body = json.loads(responses.calls[0].request.body) + self.assertEqual(request_body, {"amount": "100"}) + + @responses.activate + def test_payment_refund_with_amount_and_options(self): + result = mock_file('fake_refund') + url = '{}/{}/refund'.format(self.base_url, self.payment_id) + responses.add(responses.POST, url, status=200, body=json.dumps(result), + match_querystring=True) + options = {'speed': 'normal', 'receipt': '#rec_1'} + self.assertEqual( + self.client.payment.refund(self.payment_id, 2000, options), + result + ) + request_body = json.loads(responses.calls[0].request.body) + self.assertEqual(request_body, {'amount': 2000, 'speed': 'normal', 'receipt': '#rec_1'}) + # Ensure caller's dictionary was not mutated in-place + self.assertEqual(options, {'speed': 'normal', 'receipt': '#rec_1'}) + + @responses.activate + def test_payment_refund_with_keyword_amount(self): + result = mock_file('fake_refund') + url = '{}/{}/refund'.format(self.base_url, self.payment_id) + responses.add(responses.POST, url, status=200, body=json.dumps(result), + match_querystring=True) + self.assertEqual( + self.client.payment.refund(self.payment_id, amount=3000), + result + ) + request_body = json.loads(responses.calls[0].request.body) + self.assertEqual(request_body, {'amount': 3000}) + + @responses.activate + def test_payment_refund_with_keyword_amount_and_data(self): + result = mock_file('fake_refund') + url = '{}/{}/refund'.format(self.base_url, self.payment_id) + responses.add(responses.POST, url, status=200, body=json.dumps(result), + match_querystring=True) + self.assertEqual( + self.client.payment.refund(self.payment_id, amount=3000, data={'speed': 'optimum'}), + result + ) + request_body = json.loads(responses.calls[0].request.body) + self.assertEqual(request_body, {'amount': 3000, 'speed': 'optimum'}) + + @responses.activate + def test_payment_refund_without_arguments(self): + result = mock_file('fake_refund') + url = '{}/{}/refund'.format(self.base_url, self.payment_id) + responses.add(responses.POST, url, status=200, body=json.dumps(result), + match_querystring=True) + self.assertEqual( + self.client.payment.refund(self.payment_id), + result + ) + request_body = json.loads(responses.calls[0].request.body) + self.assertEqual(request_body, {}) @responses.activate def test_payment_fetch_multiple_refund(self):