# archive.py -- Creating an archive from a tarball
# Copyright (C) 2015 Jonas Haag <jonas@lophus.org>
# Copyright (C) 2015 Jelmer Vernooij <jelmer@jelmer.uk>
#
# SPDX-License-Identifier: Apache-2.0 OR GPL-2.0-or-later
# Dulwich is dual-licensed under the Apache License, Version 2.0 and the GNU
# General Public License as published by the Free Software Foundation; version 2.0
# or (at your option) any later version. You can redistribute it and/or
# modify it under the terms of either of these two licenses.
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.
#
# You should have received a copy of the licenses; if not, see
# <http://www.gnu.org/licenses/> for a copy of the GNU General Public License
# and <http://www.apache.org/licenses/LICENSE-2.0> for a copy of the Apache
# License, Version 2.0.
#

"""Generates tarballs for Git trees."""

__all__ = ["ChunkedBytesIO", "UnsafeArchivePathError", "tar_stream"]

import posixpath
import stat
import struct
import tarfile
from collections.abc import Generator
from contextlib import closing
from io import BytesIO
from os import SEEK_END
from typing import TYPE_CHECKING, Any, cast

if TYPE_CHECKING:
    from .object_store import BaseObjectStore
    from .objects import TreeEntry

from .index import cleanup_mode
from .objects import Blob, Tree


class UnsafeArchivePathError(ValueError):
    """Raised when a tree contains a path that is unsafe to include in an archive.

    Mirrors git's ``verify_path`` rejection (``error: invalid path``) for tree
    entries whose names would escape the archive root or alias a special
    directory on extraction. Examples include absolute paths, ``..`` or
    ``.git`` components, embedded backslashes, and NTFS alternate-data-stream
    separators (``:``).
    """

    def __init__(self, path: bytes) -> None:
        """Initialize the exception with the offending path."""
        self.path = path
        super().__init__(
            f"unsafe path in tree: {path!r}",
        )


# Characters that are never safe in a tree path destined for an archive:
# NUL terminates C strings, ``\`` is the Windows path separator (and would
# let a hostile name like ``..\evil.txt`` slip past a POSIX-only split), and
# ``:`` is the NTFS alternate-data-stream / drive-letter separator.
_UNSAFE_PATH_CHARS = (b"\x00", b"\\", b":")

_INVALID_PATH_COMPONENTS = (b".", b"..", b".git", b"")


def _is_unsafe_archive_path(path: bytes) -> bool:
    """Return True if ``path`` should not be emitted into an archive.

    Matches the spirit of git's ``verify_path``: absolute paths, ``.``/``..``/
    ``.git`` components (case-insensitive), embedded NUL bytes, backslashes,
    and colons are all rejected, regardless of host platform, because the
    resulting archive may be extracted anywhere.
    """
    if not path:
        return True
    for ch in _UNSAFE_PATH_CHARS:
        if ch in path:
            return True
    for component in path.split(b"/"):
        # Match git's protectNTFS-style normalization: trailing dots and
        # spaces are stripped before comparison so ``.git.`` / ``.GIT `` are
        # still recognised as ``.git`` on Windows.
        normalized = component.lower().rstrip(b". ")
        if normalized in _INVALID_PATH_COMPONENTS:
            return True
    return False


class ChunkedBytesIO:
    """Turn a list of bytestrings into a file-like object.

    This is similar to creating a `BytesIO` from a concatenation of the
    bytestring list, but saves memory by NOT creating one giant bytestring
    first::

        BytesIO(b''.join(list_of_bytestrings)) =~= ChunkedBytesIO(
            list_of_bytestrings)
    """

    def __init__(self, contents: list[bytes]) -> None:
        """Initialize ChunkedBytesIO.

        Args:
            contents: List of byte chunks
        """
        self.contents = contents
        self.pos = (0, 0)

    def read(self, maxbytes: int | None = None) -> bytes:
        """Read bytes from the chunked stream.

        Args:
            maxbytes: Maximum number of bytes to read (None for all)

        Returns:
            Bytes read
        """
        if maxbytes is None or maxbytes < 0:
            remaining = None
        else:
            remaining = maxbytes

        buf = []
        chunk, cursor = self.pos

        while chunk < len(self.contents):
            chunk_remainder = len(self.contents[chunk]) - cursor
            if remaining is not None and remaining < chunk_remainder:
                buf.append(self.contents[chunk][cursor : cursor + remaining])
                cursor += remaining
                self.pos = (chunk, cursor)
                break
            else:
                buf.append(self.contents[chunk][cursor:])
                if remaining is not None:
                    remaining -= chunk_remainder
                chunk += 1
                cursor = 0
                self.pos = (chunk, cursor)
        return b"".join(buf)


def tar_stream(
    store: "BaseObjectStore",
    tree: "Tree",
    mtime: int,
    prefix: bytes = b"",
    format: str = "",
) -> Generator[bytes, None, None]:
    """Generate a tar stream for the contents of a Git tree.

    Returns a generator that lazily assembles a .tar.gz archive, yielding it in
    pieces (bytestrings). To obtain the complete .tar.gz binary file, simply
    concatenate these chunks.

    Args:
      store: Object store to retrieve objects from
      tree: Tree object for the tree root
      mtime: UNIX timestamp that is assigned as the modification time for
        all files, and the gzip header modification time if format='gz'
      prefix: Optional prefix to prepend to all paths in the archive
      format: Optional compression format for tarball
    Returns:
      Bytestrings
    """
    buf = BytesIO()
    mode = "w:" + format if format else "w"

    # The tarfile.open overloads are complex; cast to Any to avoid issues
    with closing(cast(Any, tarfile.open)(name=None, mode=mode, fileobj=buf)) as tar:
        if format == "gz":
            # Manually correct the gzip header file modification time so that
            # archives created from the same Git tree are always identical.
            # The gzip header file modification time is not currently
            # accessible from the tarfile API, see:
            # https://bugs.python.org/issue31526
            buf.seek(0)
            assert buf.read(2) == b"\x1f\x8b", "Invalid gzip header"
            buf.seek(4)
            buf.write(struct.pack("<L", mtime))
            buf.seek(0, SEEK_END)

        for entry_abspath, entry in _walk_tree(store, tree, prefix):
            assert entry.sha is not None
            try:
                blob = store[entry.sha]
            except KeyError:
                # Entry probably refers to a submodule, which we don't yet
                # support.
                continue
            if isinstance(blob, Blob):
                data = ChunkedBytesIO(blob.chunked)
            else:
                # Fallback for objects without chunked attribute
                data = ChunkedBytesIO([blob.as_raw_string()])

            info = tarfile.TarInfo()
            # tarfile only works with ascii.
            info.name = entry_abspath.decode("utf-8", "surrogateescape")
            info.size = blob.raw_length()
            assert entry.mode is not None
            # Canonicalize the tree-supplied mode the same way git's archive
            # writer does before emitting a permission field. A crafted tree
            # entry can carry setuid/setgid/sticky or other non-canonical bits
            # (e.g. 0o104755), which tarfile would otherwise copy verbatim into
            # the archive so an extracted file lands setuid. cleanup_mode maps a
            # regular file to 0o644/0o755 (matching checkout), and masking to the
            # permission bits drops the special bits.
            info.mode = cleanup_mode(entry.mode) & 0o777
            info.mtime = mtime

            tar.addfile(info, data)
            yield buf.getvalue()
            buf.truncate(0)
            buf.seek(0)
    yield buf.getvalue()


def _walk_tree(
    store: "BaseObjectStore", tree: "Tree", root: bytes = b""
) -> Generator[tuple[bytes, "TreeEntry"], None, None]:
    """Recursively walk a dulwich Tree, yielding tuples of (absolute path, TreeEntry) along the way.

    Raises:
      UnsafeArchivePathError: if a tree entry's name is unsafe to emit into
        an archive (matches the spirit of git's ``verify_path`` /
        ``error: invalid path``).
    """
    for entry in tree.iteritems():
        assert entry.path is not None
        if _is_unsafe_archive_path(entry.path):
            raise UnsafeArchivePathError(entry.path)
        entry_abspath = posixpath.join(root, entry.path)
        assert entry.mode is not None
        if stat.S_ISDIR(entry.mode):
            assert entry.sha is not None
            subtree = store[entry.sha]
            if isinstance(subtree, Tree):
                yield from _walk_tree(store, subtree, entry_abspath)
        else:
            yield (entry_abspath, entry)
