diff --git a/src/python_workflow_submitter/check_visit.py b/src/python_workflow_submitter/check_visit.py new file mode 100644 index 0000000..03d624d --- /dev/null +++ b/src/python_workflow_submitter/check_visit.py @@ -0,0 +1,5 @@ +import re + + +def check_visit(vis: str): + return re.match(r"[a-zA-Z][a-zA-Z]\d\d\d\d\d-\d", vis) diff --git a/src/python_workflow_submitter/list_workflows.py b/src/python_workflow_submitter/list_workflows.py index 9daba44..a3d069d 100644 --- a/src/python_workflow_submitter/list_workflows.py +++ b/src/python_workflow_submitter/list_workflows.py @@ -6,8 +6,7 @@ from gql.transport.aiohttp import AIOHTTPTransport from python_workflow_submitter.auth.keycloak_checker import set_token_env_variable - -# TODO check visit is correct (regex) +from python_workflow_submitter.check_visit import check_visit # TODO add in maintainer when fixed @@ -124,20 +123,23 @@ async def list_workflows_in_visit( } """) - result = await client.execute_async( - query, - variable_values={ - "visit": { - "proposalCode": visit[:2], - "proposalNumber": int(visit[2:7]), - "number": int(visit[-1]), + if check_visit(visit): + result = await client.execute_async( + query, + variable_values={ + "visit": { + "proposalCode": visit[:2], + "proposalNumber": int(visit[2:7]), + "number": int(visit[-1]), + }, + "limit": limit, + "filter": filter, }, - "limit": limit, - "filter": filter, - }, - ) - json_result = json.loads(re.sub(r"'", '"', str(result))) - print(json.dumps(json_result, indent=2)) + ) + json_result = json.loads(re.sub(r"'", '"', str(result))) + print(json.dumps(json_result, indent=2)) + else: + print(f"Visit '{visit}' is invalid.") async def info_about_workflow( @@ -152,45 +154,48 @@ async def info_about_workflow( visit (str, optional): The visit the workflow is running / was ran on. Defaults to str(os.environ.get("VISIT")). """ - token: str = set_token_env_variable() - transport = AIOHTTPTransport( - url="https://workflows.diamond.ac.uk/graphql", - headers={"Authorization": f"Bearer {token}"}, - ) - client = Client( - transport=transport, - fetch_schema_from_transport=True, - ) - query = gql(""" - query Workflow( - $visit: VisitInput!, - $name: String!, - ) { - workflow( - visit: $visit, - name: $name, - ) { - name - parameters - templateRef - creator {creatorId} - status {__typename} + if check_visit(visit): + token: str = set_token_env_variable() + transport = AIOHTTPTransport( + url="https://workflows.diamond.ac.uk/graphql", + headers={"Authorization": f"Bearer {token}"}, + ) + client = Client( + transport=transport, + fetch_schema_from_transport=True, + ) + query = gql(""" + query Workflow( + $visit: VisitInput!, + $name: String!, + ) { + workflow( + visit: $visit, + name: $name, + ) { + name + parameters + templateRef + creator {creatorId} + status {__typename} + } } -} - """) + """) - result = await client.execute_async( - query, - variable_values={ - "visit": { - "proposalCode": visit[:2], - "proposalNumber": int(visit[2:7]), - "number": int(visit[-1]), + result = await client.execute_async( + query, + variable_values={ + "visit": { + "proposalCode": visit[:2], + "proposalNumber": int(visit[2:7]), + "number": int(visit[-1]), + }, + "name": name, }, - "name": name, - }, - ) - if result["workflow"] is not None: - print(json.dumps(result, indent=2)) + ) + if result["workflow"] is not None: + print(json.dumps(result, indent=2)) + else: + print(f"No workflow with name {name} found.") else: - print(f"No workflow with name {name} found.") + print(f"Visit '{visit}' is invalid.") diff --git a/src/python_workflow_submitter/submit_workflow.py b/src/python_workflow_submitter/submit_workflow.py index 7ea70fe..4aa4bda 100644 --- a/src/python_workflow_submitter/submit_workflow.py +++ b/src/python_workflow_submitter/submit_workflow.py @@ -5,6 +5,7 @@ from gql.transport.aiohttp import AIOHTTPTransport from python_workflow_submitter.auth.keycloak_checker import set_token_env_variable +from python_workflow_submitter.check_visit import check_visit from python_workflow_submitter.lintyaml import lint_yaml @@ -19,12 +20,66 @@ async def submit_workflow_yaml( visit (str, optional): The visit to run the yaml within. Defaults to str(os.environ.get("VISIT")). """ - if lint_yaml(path): - with open(f"{path}") as yamlfile: - yamlstr = yamlfile.read().rstrip() - dotenv.load_dotenv(dotenv_path="src/.env", override=True) - token: str = set_token_env_variable() + if check_visit(visit): + if lint_yaml(path): + with open(f"{path}") as yamlfile: + yamlstr = yamlfile.read().rstrip() + dotenv.load_dotenv(dotenv_path="src/.env", override=True) + token: str = set_token_env_variable() + + transport = AIOHTTPTransport( + url="https://workflows.diamond.ac.uk/graphql", + headers={"Authorization": f"Bearer {token}"}, + ) + client = Client( + transport=transport, + fetch_schema_from_transport=True, + ) + mutation = gql(""" + mutation Submit($visit: VisitInput!, $manifest: String!) { + submitWorkflow( + visit: $visit + manifest: $manifest + ) { + name + } + } + """) + result = await client.execute_async( + mutation, + variable_values={ + "visit": { + "proposalCode": str(visit[:2]), + "proposalNumber": int(visit[2:7]), + "number": int(visit[-1]), + }, + "manifest": f"""{yamlstr}""", + }, + ) + name = str(result["submitWorkflow"]["name"]) + print(f"Job '{name}' submitted to {visit}") + else: + print("Yaml did not successfully lint, not submitting.") + else: + print(f"Visit '{visit}' is invalid.") + + +async def submit_workflow( + name: str, + parameters: dict, + visit: str = str(os.environ.get("VISIT")), +): + """Submits a clusterWorkflowTemplate already in the platform. + + Args: + name (str): Name of the workflow you wish to run. + parameters (dict): paramaters to parse into the workflow. + visit (str, optional): The visit to run the yaml within. + Defaults to str(os.environ.get("VISIT")). + """ + if check_visit(visit): + token: str = set_token_env_variable() transport = AIOHTTPTransport( url="https://workflows.diamond.ac.uk/graphql", headers={"Authorization": f"Bearer {token}"}, @@ -34,77 +89,30 @@ async def submit_workflow_yaml( fetch_schema_from_transport=True, ) mutation = gql(""" - mutation Submit($visit: VisitInput!, $manifest: String!) { - submitWorkflow( + mutation SubmitGeneric + ($name: String!, $visit: VisitInput!, $parameters: JSON!){ + submitWorkflowTemplate( + name: $name visit: $visit - manifest: $manifest - ) { + parameters: $parameters + ){ name - } + } } """) result = await client.execute_async( mutation, variable_values={ + "name": name, "visit": { "proposalCode": str(visit[:2]), "proposalNumber": int(visit[2:7]), "number": int(visit[-1]), }, - "manifest": f"""{yamlstr}""", + "parameters": parameters, }, ) - name = str(result["submitWorkflow"]["name"]) + name = str(result["submitWorkflowTemplate"]["name"]) print(f"Job '{name}' submitted to {visit}") else: - print("Yaml did not successfully lint, not submitting.") - - -async def submit_workflow( - name: str, - parameters: dict, - visit: str = str(os.environ.get("VISIT")), -): - """Submits a clusterWorkflowTemplate already in the platform. - - Args: - name (str): Name of the workflow you wish to run. - parameters (dict): paramaters to parse into the workflow. - visit (str, optional): The visit to run the yaml within. - Defaults to str(os.environ.get("VISIT")). - """ - token: str = set_token_env_variable() - transport = AIOHTTPTransport( - url="https://workflows.diamond.ac.uk/graphql", - headers={"Authorization": f"Bearer {token}"}, - ) - client = Client( - transport=transport, - fetch_schema_from_transport=True, - ) - mutation = gql(""" -mutation SubmitGeneric -($name: String!, $visit: VisitInput!, $parameters: JSON!){ -submitWorkflowTemplate( - name: $name - visit: $visit - parameters: $parameters - ){ - name - } -} -""") - result = await client.execute_async( - mutation, - variable_values={ - "name": name, - "visit": { - "proposalCode": str(visit[:2]), - "proposalNumber": int(visit[2:7]), - "number": int(visit[-1]), - }, - "parameters": parameters, - }, - ) - name = str(result["submitWorkflowTemplate"]["name"]) - print(f"Job '{name}' submitted to {visit}") + print(f"Visit '{visit}' is invalid.") diff --git a/tests/test_list_workflows.py b/tests/test_list_workflows.py index 9ab0985..a92e475 100644 --- a/tests/test_list_workflows.py +++ b/tests/test_list_workflows.py @@ -80,6 +80,60 @@ async def test_list_workflows_in_visit( mock_json.assert_called_once() +@pytest.mark.asyncio +@patch("python_workflow_submitter.list_workflows.json.dumps") +@patch("python_workflow_submitter.list_workflows.set_token_env_variable") +@patch("python_workflow_submitter.list_workflows.Client") +@patch("python_workflow_submitter.list_workflows.print") +async def test_list_workflows_in_bad_visit( + mock_print: MagicMock, + mock_client: AsyncMock, + mock_key: MagicMock, + mock_json: MagicMock, +): + mock_instance = AsyncMock() + mock_key.return_value = "token" + mock_client.return_value = mock_instance + mock_instance.execute_async = AsyncMock( + return_value={"workflowTemplates": {"name": "workflow123"}} + ) + await list_workflows_in_visit( + limit=5, + filter={ + "creator": "gmg29649", + "template": "example-template", + "workflowStatusFilter": {"succeeded": True}, + }, + visit="BAD", + ) + mock_instance.execute_async.assert_not_called() + mock_json.assert_not_called() + mock_print.assert_called_once() + + +@pytest.mark.asyncio +@patch("python_workflow_submitter.list_workflows.print") +@patch("python_workflow_submitter.list_workflows.json.dumps") +@patch("python_workflow_submitter.list_workflows.set_token_env_variable") +@patch("python_workflow_submitter.list_workflows.Client") +async def test_info_about_workflow_bad_visit( + mock_client: AsyncMock, + mock_key: MagicMock, + mock_json: MagicMock, + mock_print: MagicMock, +): + mock_instance = AsyncMock() + mock_key.return_value = "token" + mock_client.return_value = mock_instance + mock_instance.execute_async = AsyncMock( + return_value={"workflow": {"name": "workflow123"}} + ) + await info_about_workflow(name="fakename", visit="BAD") + mock_instance.execute_async.assert_not_called() + mock_json.assert_not_called() + mock_print.assert_called_once() + + @pytest.mark.asyncio @patch("python_workflow_submitter.list_workflows.json.dumps") @patch("python_workflow_submitter.list_workflows.set_token_env_variable") diff --git a/tests/test_submit_to_graphql.py b/tests/test_submit_to_graphql.py index e7f7f5c..098820e 100644 --- a/tests/test_submit_to_graphql.py +++ b/tests/test_submit_to_graphql.py @@ -76,3 +76,51 @@ async def test_submit_stock_workflow( ) await submit_workflow(name="workflow123", parameters={}, visit="ks10000-3") mock_instance.execute_async.assert_called_once() + + +@pytest.mark.asyncio +@patch("python_workflow_submitter.submit_workflow.print") +@patch("python_workflow_submitter.submit_workflow.dotenv.load_dotenv") +@patch("python_workflow_submitter.submit_workflow.open") +@patch("python_workflow_submitter.submit_workflow.lint_yaml") +@patch("python_workflow_submitter.submit_workflow.set_token_env_variable") +@patch("python_workflow_submitter.submit_workflow.Client") +async def test_submit_workflow_to_graphql_bad_visit( + mock_client: AsyncMock, + mock_key: MagicMock, + mock_lint: MagicMock, + mock_open: MagicMock, + mock_load_env: MagicMock, + mock_print: MagicMock, +): + mock_lint.return_value = True + mock_instance = AsyncMock() + mock_key.return_value = "token" + mock_client.return_value = mock_instance + mock_instance.execute_async = AsyncMock( + return_value={"submitWorkflow": {"name": "workflow123"}} + ) + await submit_workflow_yaml("fakeyaml", visit="BAD") + mock_load_env.assert_not_called() + mock_open.assert_not_called() + mock_print.assert_called_once() + + +@pytest.mark.asyncio +@patch("python_workflow_submitter.submit_workflow.print") +@patch("python_workflow_submitter.submit_workflow.set_token_env_variable") +@patch("python_workflow_submitter.submit_workflow.Client") +async def test_submit_stock_workflow_bad_visit( + mock_client: AsyncMock, + mock_key: MagicMock, + mock_print: MagicMock, +): + + mock_instance = AsyncMock() + mock_key.return_value = "token" + mock_client.return_value = mock_instance + mock_instance.execute_async = AsyncMock( + return_value={"submitWorkflowTemplate": {"name": "workflow123"}} + ) + await submit_workflow(name="workflow123", parameters={}, visit="BAD") + mock_print.assert_called_once()