From 4dd78d66d26c023db6c12f469a11529fb1e3e993 Mon Sep 17 00:00:00 2001 From: kelly Date: Mon, 17 Aug 2026 17:28:46 +1000 Subject: [PATCH] fix: preserve the indentation characters of the original code Use the indentation characters of the docstring's opening as the indentation for the lines inside the docstring, instead of replacing them with spaces. Similarly, the use of `tokenize.untokenize` needs to preserve the whitespace that was used, to avoid changing the indentation of the original code. --- src/docformatter/format.py | 49 ++++++++++++++++++-- tests/_data/string_files/do_format_code.toml | 21 +++++++++ tests/formatter/test_do_format_code.py | 2 + 3 files changed, 69 insertions(+), 3 deletions(-) diff --git a/src/docformatter/format.py b/src/docformatter/format.py index 584ff6a..03ba1e0 100644 --- a/src/docformatter/format.py +++ b/src/docformatter/format.py @@ -575,6 +575,45 @@ def _get_unmatched_start_end_indices( return (_start_row, _start_col), (_end_row, _end_col) +class _Untokenizer(tokenize.Untokenizer): + __line = "" + + def untokenize(self, iterable): + def _remember_line(tokens): + for tok in tokens: + self.__prev_line, self.__line = self.__line, tok[4] + yield tok + + return super().untokenize(_remember_line(iterable)) + + def add_backslash_continuation(self, start): + """Add backslash continuation characters if the row has increased + without encountering a newline token. + + This also inserts the correct amount of whitespace before the backslash. + """ + row_offset = start[0] - self.prev_row + if row_offset == 0: + return + + newline = "\r\n" if self.__prev_line.endswith("\r\n") else "\n" + line = self.__prev_line.rstrip("\\\r\n") + ws = line[len(line.rstrip()):] + self.tokens.append(ws + f"\\{newline}" * row_offset) + self.prev_col = 0 + + def add_whitespace(self, start, line=""): + row, col = start + if row < self.prev_row or row == self.prev_row and col < self.prev_col: + raise ValueError("start ({},{}) precedes previous end ({},{})" + .format(row, col, self.prev_row, self.prev_col)) + self.add_backslash_continuation(start) + col_offset = col - self.prev_col + if col_offset: + line = line or self.__line + self.tokens.append(line[self.prev_col : col]) + + class FormatResult: """Possible exit codes.""" @@ -753,7 +792,9 @@ def _do_add_formatted_docstring( blank_line_count : int The number of blank lines to add after the docstring. """ - _indent = " " * token.start[1] if docstring_type != "module" else "" + _indent = ( + token.line[: token.start[1]] if docstring_type != "module" else "" + ) _formatted = self._do_format_docstring(_indent, token.string) _line = _indent + _formatted @@ -817,7 +858,9 @@ def _do_add_unformatted_docstring( docstring_type : str The type of the docstring (e.g., module, class, function, attribute). """ - _indent = " " * token.start[1] if docstring_type != "module" else "" + _indent = ( + token.line[: token.start[1]] if docstring_type != "module" else "" + ) _line = _indent + token.string _new_token = tokenize.TokenInfo( type=tokenize.STRING, @@ -916,7 +959,7 @@ def _do_format_code(self, source: str) -> str: # Perform docstring rewriting self._do_rewrite_docstring_blocks(tokens) - _code = tokenize.untokenize(self.new_tokens) + _code = _Untokenizer().untokenize(self.new_tokens) return _strings.do_normalize_line_endings( _code.splitlines(True), _original_newline diff --git a/tests/_data/string_files/do_format_code.toml b/tests/_data/string_files/do_format_code.toml index d652a15..ceaa402 100644 --- a/tests/_data/string_files/do_format_code.toml +++ b/tests/_data/string_files/do_format_code.toml @@ -170,6 +170,27 @@ expected='''def foo(): if True: x = 1''' +[wrapped_indentation] +source='''def foo(): + """This is a very, very, very long docstring that should really be reformatted nicely by docformatter.""" + if True: + x = 1''' +expected='''def foo(): + """This is a very, very, very long docstring that should really be reformatted + nicely by docformatter.""" + if True: + x = 1''' + +[preserve_whitespace] +source='''def foo( + bar +): + pass''' +expected='''def foo( + bar +): + pass''' + [escaped_newlines] source='''def foo(): """ diff --git a/tests/formatter/test_do_format_code.py b/tests/formatter/test_do_format_code.py index 037ace1..a0adef4 100644 --- a/tests/formatter/test_do_format_code.py +++ b/tests/formatter/test_do_format_code.py @@ -72,6 +72,8 @@ ("non_docstring", NO_ARGS), ("tabbed_indentation", NO_ARGS), ("mixed_indentation", NO_ARGS), + ("wrapped_indentation", NO_ARGS), + ("preserve_whitespace", NO_ARGS), ("escaped_newlines", NO_ARGS), ("code_comments", NO_ARGS), ("inline_comment", NO_ARGS),