Skip to content
Merged
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
12 changes: 6 additions & 6 deletions src/scriptworker/context.py
Original file line number Diff line number Diff line change
Expand Up @@ -59,7 +59,7 @@ class Context(object):

config: Optional[Dict[str, Any]] = None
credentials_timestamp: Optional[int] = None
credentials_fd: int = -1
credentials_file = None
proc: Optional[task_process.TaskProcess] = None
queue: Optional[Queue] = None
session: Optional[aiohttp.ClientSession] = None
Expand Down Expand Up @@ -99,17 +99,17 @@ def claim_task(self, claim_task: Optional[Dict[str, Any]]) -> None:
if claim_task:
self.task = claim_task["task"]
self.verify_task()
# flags=0 to let the child inherit this fd
self.credentials_fd = os.memfd_create("scriptworker_temp_creds", flags=0)
self.credentials_file = tempfile.TemporaryFile()
self.temp_credentials = claim_task["credentials"]
path = os.path.join(self.config["work_dir"], "task.json")
assert self.task
self.write_json(path, self.task, "Writing task file to {path}...")
else:
self.temp_credentials = None
self.task = None
os.close(self.credentials_fd)
self.credentials_fd = -1
if self.credentials_file is not None:
self.credentials_file.close()
self.credentials_file = None

def verify_task(self) -> None:
"""Run some task sanity checks on ``self.task``."""
Expand Down Expand Up @@ -201,7 +201,7 @@ def temp_credentials(self, credentials: Optional[Dict[str, Any]]) -> None:
if credentials:
data = json.dumps(credentials, indent=2, sort_keys=True).encode("ascii")
# use pwrite so we don't confuse the child by changing the file offset
assert os.pwrite(self.credentials_fd, data, 0) == len(data)
assert os.pwrite(self.credentials_file.fileno(), data, 0) == len(data)

def write_json(self, path: str, contents: Dict[str, Any], message: str) -> None:
"""Write json to disk.
Expand Down
5 changes: 3 additions & 2 deletions src/scriptworker/task.py
Original file line number Diff line number Diff line change
Expand Up @@ -672,15 +672,16 @@ async def run_task(context, to_cancellable_process):
env["TASK_ID"] = context.task_id or "None"
env["RUN_ID"] = str(get_run_id(context.claim_task))
env["TASKCLUSTER_ROOT_URL"] = context.config["taskcluster_root_url"]
env["TASKCLUSTER_CREDENTIALS_FD"] = str(context.credentials_fd)
credentials_fd = context.credentials_file.fileno()
env["TASKCLUSTER_CREDENTIALS_FD"] = str(credentials_fd)
kwargs = {
"stdout": PIPE,
"stderr": PIPE,
"stdin": None,
"close_fds": True,
"preexec_fn": lambda: os.setsid(),
"env": env,
"pass_fds": (context.credentials_fd,),
"pass_fds": (credentials_fd,),
} # pragma: no branch

timeout = get_task_maxruntime(context.task, context.config["task_max_timeout"])
Expand Down
24 changes: 12 additions & 12 deletions tests/test_context.py
Original file line number Diff line number Diff line change
Expand Up @@ -77,42 +77,42 @@ async def test_set_reset_task(rw_context, claim_task, reclaim_task):
assert rw_context.temp_queue is None


def test_credentials_fd_initial(rw_context):
assert rw_context.credentials_fd == -1
def test_credentials_file_initial(rw_context):
assert rw_context.credentials_file is None


@pytest.mark.asyncio
async def test_credentials_fd_opened_on_claim_task(rw_context, claim_task):
async def test_credentials_file_opened_on_claim_task(rw_context, claim_task):
rw_context.claim_task = claim_task
assert rw_context.credentials_fd >= 0
os.fstat(rw_context.credentials_fd) # raises OSError if fd is invalid
assert rw_context.credentials_file is not None
os.fstat(rw_context.credentials_file.fileno()) # raises OSError if fd is invalid


@pytest.mark.asyncio
async def test_credentials_fd_content(rw_context, claim_task):
async def test_credentials_file_content(rw_context, claim_task):
rw_context.claim_task = claim_task
fd = rw_context.credentials_fd
fd = rw_context.credentials_file.fileno()
size = os.fstat(fd).st_size
data = os.pread(fd, size, 0)
assert json.loads(data) == claim_task["credentials"]


@pytest.mark.asyncio
async def test_credentials_fd_updated_on_reclaim(rw_context, claim_task, reclaim_task):
async def test_credentials_file_updated_on_reclaim(rw_context, claim_task, reclaim_task):
rw_context.claim_task = claim_task
rw_context.reclaim_task = reclaim_task
fd = rw_context.credentials_fd
fd = rw_context.credentials_file.fileno()
size = os.fstat(fd).st_size
data = os.pread(fd, size, 0)
assert json.loads(data) == reclaim_task["credentials"]


@pytest.mark.asyncio
async def test_credentials_fd_closed_on_reset(rw_context, claim_task):
async def test_credentials_file_closed_on_reset(rw_context, claim_task):
rw_context.claim_task = claim_task
fd = rw_context.credentials_fd
fd = rw_context.credentials_file.fileno()
rw_context.claim_task = None
assert rw_context.credentials_fd == -1
assert rw_context.credentials_file is None
with pytest.raises(OSError):
os.fstat(fd)

Expand Down