Skip to content

Commit 7181e6e

Browse files
authored
Merge pull request #355 from avinxshKD/fix/server-toctou-races
fix: server workflow API TOCTOU races and misc issues
2 parents 3ce7520 + f5d8655 commit 7181e6e

4 files changed

Lines changed: 61 additions & 20 deletions

File tree

server/controller/workflow.py

Lines changed: 2 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -30,7 +30,7 @@ def getAllActionHash(root):
3030
def postWorkflow():
3131
try:
3232
lastestHash = getLasteshActionHash(ET.fromstring(request.data))
33-
except:
33+
except Exception:
3434
return "Invalid GraphML", 400
3535
graphML = request.data.decode('utf')
3636
return workFlowModel.insert(graphML, lastestHash)
@@ -60,7 +60,7 @@ def updateWorkflow(serverID):
6060
latestHash = getLasteshActionHash(root)
6161
if(not forceUpdate):
6262
allHash = getAllActionHash(root)
63-
except:
63+
except Exception:
6464
return "Invalid GraphML", 400
6565
graphML = request.data.decode('utf')
6666
if(forceUpdate):

server/model/workflows.py

Lines changed: 18 additions & 18 deletions
Original file line numberDiff line numberDiff line change
@@ -1,11 +1,11 @@
11
from pymongo import MongoClient
2-
from pymongo import MongoClient
2+
from pymongo.errors import DuplicateKeyError
33
import time
44
from bson.objectid import ObjectId
55
from bson.errors import InvalidId
66
import os
77
import xml.etree.ElementTree as ET
8-
import random
8+
import secrets
99
import string
1010
from dotenv import load_dotenv
1111
load_dotenv()
@@ -14,20 +14,21 @@ class WorkFlowModel:
1414
def __init__(self) -> None:
1515
self.collection = MongoClient(os.getenv('MongoURL'))[
1616
os.getenv('dbName')][os.getenv('tableName')]
17+
self.collection.create_index('serverID', unique=True)
1718

1819
def get_random_string(self, length):
1920
letters = string.ascii_letters+string.digits
20-
return ''.join(random.choice(letters) for i in range(length))
21+
return ''.join(secrets.choice(letters) for i in range(length))
2122

2223
def insert(self, graphml, latestHash):
23-
serverID = ""
2424
while(True):
2525
serverID = self.get_random_string(6)
26-
if(not self.collection.find_one({'serverID': serverID})):
27-
break
28-
self.collection.insert_one(
29-
{'graphml': graphml, 'latestHash': latestHash, 'serverID': serverID})
30-
return serverID
26+
try:
27+
self.collection.insert_one(
28+
{'graphml': graphml, 'latestHash': latestHash, 'serverID': serverID})
29+
return serverID
30+
except DuplicateKeyError:
31+
continue
3132

3233
def get(self, serverID):
3334
cl = self.collection.find_one({'serverID': serverID})
@@ -36,20 +37,19 @@ def get(self, serverID):
3637
return cl['graphml']
3738

3839
def update(self, serverID, graphml, latestHash, allHash):
39-
existingRecord = self.collection.find_one({'serverID': serverID})
40+
existingRecord = self.collection.find_one_and_update(
41+
{'serverID': serverID, 'latestHash': {'$in': allHash}},
42+
{"$set": {'graphml': graphml, 'latestHash': latestHash}})
4043
if existingRecord is None:
41-
return False, 'serverID do not exists.'
42-
latestExistingHash = existingRecord['latestHash']
43-
if latestExistingHash not in allHash:
44+
if not self.collection.find_one({'serverID': serverID}):
45+
return False, 'serverID do not exists.'
4446
return False, 'Can not update as provided graph do not has latest changes.'
45-
self.collection.update_one({'serverID': serverID}, {
46-
"$set": {'graphml': graphml, 'latestHash': latestHash}})
4747
return True, latestHash
4848

4949
def forceUpdate(self, serverID, graphml, latestHash):
50-
existingRecord = self.collection.find_one({'serverID': serverID})
50+
existingRecord = self.collection.find_one_and_update(
51+
{'serverID': serverID},
52+
{"$set": {'graphml': graphml, 'latestHash': latestHash}})
5153
if existingRecord is None:
5254
return False, 'serverID do not exists.'
53-
self.collection.update_one({'serverID': serverID},
54-
{"$set": {'graphml': graphml, 'latestHash': latestHash}})
5555
return True, latestHash

server/server.py

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -6,6 +6,7 @@
66
load_dotenv()
77

88
app = Flask(__name__)
9+
app.config['MAX_CONTENT_LENGTH'] = 2 * 1024 * 1024
910
CORS(app)
1011

1112
app.register_blueprint(workFlow, url_prefix='/workflow')

server/tests/test_workflow_controller.py

Lines changed: 40 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -22,6 +22,15 @@ def __init__(self, graph_response):
2222
def get(self, _server_id):
2323
return self.graph_response
2424

25+
def insert(self, graphml, latestHash):
26+
return 'test01'
27+
28+
def update(self, serverID, graphml, latestHash, allHash):
29+
return (True, latestHash)
30+
31+
def forceUpdate(self, serverID, graphml, latestHash):
32+
return (True, latestHash)
33+
2534

2635
class WorkflowControllerTests(unittest.TestCase):
2736
@classmethod
@@ -77,6 +86,37 @@ def test_hash_header_returns_200_for_matching_history(self):
7786
self.assertEqual(response.status_code, 200)
7887
self.assertEqual(response.get_data(as_text=True), VALID_GRAPHML)
7988

89+
def test_post_workflow_returns_server_id(self):
90+
client = self.make_client(None)
91+
response = client.post('/workflow/', data=VALID_GRAPHML,
92+
content_type='application/xml')
93+
self.assertEqual(response.status_code, 200)
94+
self.assertEqual(response.get_data(as_text=True), 'test01')
95+
96+
def test_post_workflow_invalid_xml_returns_400(self):
97+
client = self.make_client(None)
98+
response = client.post('/workflow/', data=b'not xml',
99+
content_type='application/xml')
100+
self.assertEqual(response.status_code, 400)
101+
102+
def test_update_workflow_returns_200(self):
103+
client = self.make_client(None)
104+
response = client.post('/workflow/test01', data=VALID_GRAPHML,
105+
content_type='application/xml')
106+
self.assertEqual(response.status_code, 200)
107+
108+
def test_update_workflow_invalid_xml_returns_400(self):
109+
client = self.make_client(None)
110+
response = client.post('/workflow/test01', data=b'not xml',
111+
content_type='application/xml')
112+
self.assertEqual(response.status_code, 400)
113+
114+
def test_force_update_workflow_returns_200(self):
115+
client = self.make_client(None)
116+
response = client.post('/workflow/test01?force=true', data=VALID_GRAPHML,
117+
content_type='application/xml')
118+
self.assertEqual(response.status_code, 200)
119+
80120

81121
if __name__ == '__main__':
82122
unittest.main()

0 commit comments

Comments
 (0)