跳转至

第 1 周 参考答案(默认折叠)

做完再看

先让自己的测试全部通过,再展开对照。看完合上,凭记忆把自己的版本重写一遍——只看不写等于没学(见 ai-guide.md 规则 5)。

点击展开:project_pipeline/pipeline.py
exercises/solutions/week1/project_pipeline/pipeline.py
"""Week 1 小项目 · 文本处理管线(参考答案)

用生成器管线统计一个目录下全部 .md 文件的词频、标题数、代码块数。
用法:uv run python exercises/week1/project_pipeline/pipeline.py <目录> [--top 20]
验证:与 exercises/week1/project_pipeline/test_pipeline.py 复制到同一临时目录跑 pytest。

设计要点:
- 每个步骤都是"接收可迭代、产出迭代器"的生成器函数,用 @register_step 注册;
- 出错的文件不中断管线:Record.error 记录原因,后续步骤跳过它,summarize 统计错误数;
- Pipeline 按名字从注册表取步骤并串起来;Timer/@timer 负责计时。
"""

from __future__ import annotations

import argparse
import functools
import re
import time
from collections import Counter
from collections.abc import Callable, Iterable, Iterator
from dataclasses import dataclass, field
from pathlib import Path
from typing import Any, Protocol

# ---------- 数据模型 ----------


@dataclass
class Record:
    """管线里流动的一条记录:一个文件。"""

    path: Path
    text: str = ""
    words: list[str] = field(default_factory=list)
    headings: int = 0
    code_blocks: int = 0
    error: str | None = None


@dataclass(frozen=True)
class Stats:
    """最终统计结果(不可变)。"""

    files: int
    errors: int
    headings: int
    code_blocks: int
    top_words: tuple[tuple[str, int], ...]


class PipelineError(Exception):
    """管线内部错误的基类。"""


# ---------- 注册表 ----------


class Step(Protocol):
    """一个步骤:接收可迭代对象,返回迭代器(通常是生成器函数)。"""

    def __call__(self, items: Iterable[Any]) -> Iterator[Any]: ...


class Registry:
    """名字 → 步骤 的注册表;`@registry.register("name")` 或 `@registry.register`。"""

    def __init__(self) -> None:
        self._steps: dict[str, Step] = {}

    def register(self, name: str | Callable | None = None) -> Any:
        def decorator(fn: Step) -> Step:
            key = fn.__name__ if name is None or callable(name) else name
            if key in self._steps:
                raise ValueError(f"步骤重名:{key}")
            self._steps[key] = fn
            return fn

        if callable(name):  # 不带括号使用:@registry.register
            return decorator(name)
        return decorator

    def get(self, name: str) -> Step:
        try:
            return self._steps[name]
        except KeyError as e:
            raise PipelineError(f"未知步骤:{name}") from e

    def names(self) -> list[str]:
        return list(self._steps)


registry = Registry()
register_step = registry.register

# ---------- 计时 ----------


class Timer:
    """上下文管理器:`with Timer() as t: ...; t.elapsed`。"""

    def __enter__(self) -> Timer:
        self._start = time.perf_counter()
        self.elapsed = 0.0
        return self

    def __exit__(self, *exc: object) -> None:
        self.elapsed = time.perf_counter() - self._start


def timer[**P, R](fn: Callable[P, R]) -> Callable[P, R]:
    """装饰器:把耗时(秒)记到 fn.last_elapsed 上。"""

    @functools.wraps(fn)
    def wrapper(*args: P.args, **kwargs: P.kwargs) -> R:
        with Timer() as t:
            result = fn(*args, **kwargs)
        wrapper.last_elapsed = t.elapsed  # type: ignore[attr-defined]
        return result

    wrapper.last_elapsed = 0.0  # type: ignore[attr-defined]
    return wrapper


# ---------- 步骤 ----------

CODE_BLOCK = re.compile(r"```.*?```", re.DOTALL)
HEADING = re.compile(r"^#{1,6}\s+\S", re.MULTILINE)
WORD = re.compile(r"[A-Za-z][A-Za-z'-]+|[\u4e00-\u9fff]+")


def iter_markdown(root: Path) -> Iterator[Path]:
    """惰性遍历目录下全部 .md 文件(按路径排序保证可重复)。"""
    yield from sorted(p for p in root.rglob("*.md") if p.is_file())


@register_step("read_files")
def read_files(paths: Iterable[Path]) -> Iterator[Record]:
    """读文件成 Record;读不了的文件不中断,记入 Record.error。"""
    for path in paths:
        try:
            yield Record(path=path, text=path.read_text(encoding="utf-8"))
        except (OSError, UnicodeDecodeError) as e:
            yield Record(path=path, error=f"{type(e).__name__}: {e}")


@register_step("strip_code_blocks")
def strip_code_blocks(records: Iterable[Record]) -> Iterator[Record]:
    """统计并删除 ``` 代码块,避免代码里的词污染词频。"""
    for r in records:
        if r.error is None:
            r.code_blocks = len(CODE_BLOCK.findall(r.text))
            r.text = CODE_BLOCK.sub(" ", r.text)
        yield r


@register_step("count_headings")
def count_headings(records: Iterable[Record]) -> Iterator[Record]:
    """统计 Markdown 标题行数。"""
    for r in records:
        if r.error is None:
            r.headings = len(HEADING.findall(r.text))
        yield r


@register_step("split_words")
def split_words(records: Iterable[Record]) -> Iterator[Record]:
    """切词:英文按单词、中文按连续汉字串。"""
    for r in records:
        if r.error is None:
            r.words = WORD.findall(r.text)
        yield r


@register_step("normalize")
def normalize(records: Iterable[Record]) -> Iterator[Record]:
    """小写化并去掉 1 个字符的词。"""
    for r in records:
        if r.error is None:
            r.words = [w.lower() for w in r.words if len(w) > 1]
        yield r


def summarize(records: Iterable[Record], top: int = 20) -> Stats:
    """终点:把记录流聚合成 Stats(这一步会消耗整个迭代器)。"""
    counter: Counter[str] = Counter()
    files = errors = headings = code_blocks = 0
    for r in records:
        files += 1
        if r.error is not None:
            errors += 1
            continue
        headings += r.headings
        code_blocks += r.code_blocks
        counter.update(r.words)
    top_words = tuple(counter.most_common(top))
    return Stats(files, errors, headings, code_blocks, top_words)


# ---------- 管线 ----------

DEFAULT_STEPS = [
    "read_files",
    "strip_code_blocks",
    "count_headings",
    "split_words",
    "normalize",
]


class Pipeline:
    """按名字串起若干步骤:source → step1 → step2 → ... → summarize。"""

    def __init__(
        self, steps: Iterable[str] = DEFAULT_STEPS, registry: Registry = registry
    ) -> None:
        self.step_names = list(steps)
        self._steps = [registry.get(name) for name in self.step_names]

    def __len__(self) -> int:
        return len(self._steps)

    def __iter__(self) -> Iterator[str]:
        return iter(self.step_names)

    def __repr__(self) -> str:
        return f"Pipeline({' -> '.join(self.step_names)})"

    def stream(self, source: Iterable[Any]) -> Iterator[Any]:
        """只串联步骤、不聚合,仍是惰性的。"""
        return functools.reduce(
            lambda items, step: step(items), self._steps, iter(source)
        )

    @timer
    def run(self, source: Iterable[Any], top: int = 20) -> Stats:
        return summarize(self.stream(source), top=top)


def format_report(stats: Stats, elapsed: float) -> str:
    head = f"文件 {stats.files}(失败 {stats.errors})"
    head += f"  标题 {stats.headings}  代码块 {stats.code_blocks}"
    lines = [head, f"耗时 {elapsed * 1000:.1f} ms", "词频:"]
    lines += [f"  {w:<20}{n:>6}" for w, n in stats.top_words]
    return "\n".join(lines)


def main(argv: list[str] | None = None) -> int:
    parser = argparse.ArgumentParser(
        description="统计目录下 Markdown 的词频/标题/代码块"
    )
    parser.add_argument("root", type=Path)
    parser.add_argument("--top", type=int, default=20)
    args = parser.parse_args(argv)
    if not args.root.is_dir():
        parser.error(f"{args.root} 不是目录")
    pipeline = Pipeline()
    stats = pipeline.run(iter_markdown(args.root), top=args.top)
    elapsed: float = pipeline.run.last_elapsed  # type: ignore[attr-defined]
    print(format_report(stats, elapsed))
    return 0


if __name__ == "__main__":
    raise SystemExit(main())
点击展开:w1_01_generators.py
exercises/solutions/week1/w1_01_generators.py
"""Week 1 · M1a · 迭代协议与生成器 参考答案

对照用:先自己把 exercises/week1/w1_01_generators.py 写完再看这里。
验证:复制到临时目录并与 exercises/week1/test_w1_01_generators.py
      一起运行 pytest(答案目录本身不放 test 文件,否则同名测试会冲突)。
"""

from collections.abc import Iterable, Iterator
from itertools import islice
from pathlib import Path


def countdown(n: int) -> Iterator[int]:
    """生成器:从 n 递减产出到 0(含 0)。

    例:list(countdown(3)) -> [3, 2, 1, 0]
    """
    while n >= 0:
        yield n
        n -= 1


def read_chunks(path: Path, size: int) -> Iterator[str]:
    """惰性逐块读取 UTF-8 文本文件,每次产出至多 size 个字符。

    size <= 0 时抛 ValueError("size 必须是正整数")。
    例:文件内容 "abcde"、size=2 -> "ab", "cd", "e"
    """
    if size <= 0:
        raise ValueError("size 必须是正整数")
    with path.open(encoding="utf-8") as f:
        while True:
            chunk = f.read(size)
            if not chunk:
                return
            yield chunk


def parse_records(lines: Iterable[str]) -> Iterator[dict[str, str | int]]:
    """把 "name,score" 形式的每行解析成 {"name": str, "score": int}。

    跳过空行与以 # 开头的注释行;行号从 1 开始计。
    字段数不对或 score 不是整数 -> ValueError("第 3 行不合法: ...")。
    例:["a,1", "", "# c", "b,2"] -> {"name": "a", "score": 1}, ...
    """
    for lineno, raw in enumerate(lines, start=1):
        line = raw.strip()
        if not line or line.startswith("#"):
            continue
        parts = line.split(",")
        if len(parts) != 2:
            raise ValueError(f"第 {lineno} 行不合法: {raw!r}")
        name, score = parts[0].strip(), parts[1].strip()
        try:
            yield {"name": name, "score": int(score)}
        except ValueError as e:
            raise ValueError(f"第 {lineno} 行不合法: {raw!r}") from e


def take[T](n: int, it: Iterable[T]) -> list[T]:
    """取可迭代对象的前 n 个元素(不足则全取),用 itertools.islice。

    n < 0 时抛 ValueError("n 不能是负数")。
    例:take(2, itertools.count()) -> [0, 1]
    """
    if n < 0:
        raise ValueError("n 不能是负数")
    return list(islice(it, n))


def running_mean(nums: Iterable[float]) -> Iterator[float]:
    """逐个产出"到目前为止"的平均值。

    例:list(running_mean([2, 4, 6])) -> [2.0, 3.0, 4.0]
    """
    total = 0.0
    for index, num in enumerate(nums, start=1):
        total += num
        yield total / index


def flatten(nested: Iterable[object]) -> Iterator[object]:
    """递归展平任意层嵌套的列表/元组;字符串当作原子不拆开。

    例:list(flatten([1, [2, [3, "ab"]]])) -> [1, 2, 3, "ab"]
    """
    for item in nested:
        if isinstance(item, list | tuple):
            yield from flatten(item)
        else:
            yield item


class Countdown:
    """手写迭代器类:自己实现 __iter__ 与 __next__。

    例:list(Countdown(2)) -> [2, 1, 0];同一个实例遍历一次就耗尽。
    """

    def __init__(self, start: int) -> None:
        self.current = start

    def __iter__(self) -> Countdown:
        """迭代器自己就是可迭代对象,返回 self。"""
        return self

    def __next__(self) -> int:
        """产出当前值并递减;current < 0 时抛 StopIteration。"""
        if self.current < 0:
            raise StopIteration
        value = self.current
        self.current -= 1
        return value
点击展开:w1_02_itertools.py
exercises/solutions/week1/w1_02_itertools.py
"""Week 1 · M1b · itertools 与惰性思维 参考答案

对照用:先自己把 exercises/week1/w1_02_itertools.py 写完再看这里。
验证:复制到临时目录并与 exercises/week1/test_w1_02_itertools.py
      一起运行 pytest(答案目录本身不放 test 文件,否则同名测试会冲突)。
"""

import sys
from collections import Counter, deque
from collections.abc import Callable, Iterable, Iterator
from itertools import accumulate, batched, groupby, zip_longest
from operator import itemgetter
from typing import Any

# interleave 用的哨兵:区分"真的有这个元素"和"zip_longest 的填充值"
_MISSING = object()


def chunked[T](it: Iterable[T], size: int) -> Iterator[tuple[T, ...]]:
    """把可迭代对象按 size 切成元组,最后一块可能不满。

    size <= 0 时抛 ValueError("size 必须是正整数")。
    例:list(chunked([1, 2, 3], 2)) -> [(1, 2), (3,)]
    """
    if size <= 0:
        raise ValueError("size 必须是正整数")
    return batched(it, size, strict=False)


def window[T](it: Iterable[T], n: int) -> Iterator[tuple[T, ...]]:
    """滑动窗口:每次产出连续 n 个元素组成的元组。

    n <= 0 时抛 ValueError("n 必须是正整数");元素不足 n 个时不产出任何窗口。
    例:list(window([1, 2, 3], 2)) -> [(1, 2), (2, 3)](等价 pairwise)
    """
    if n <= 0:
        raise ValueError("n 必须是正整数")
    buffer: deque[T] = deque(maxlen=n)
    for item in it:
        buffer.append(item)
        if len(buffer) == n:
            yield tuple(buffer)


def group_by_key(
    records: Iterable[dict[str, Any]], key: str
) -> dict[str, list[dict[str, Any]]]:
    """按字典里 key 字段的值分组,返回普通 dict。

    例:group_by_key([{"t": "a"}, {"t": "b"}, {"t": "a"}], "t")
        -> {"a": [{"t": "a"}, {"t": "a"}], "b": [{"t": "b"}]}
    """
    get = itemgetter(key)
    ordered = sorted(records, key=get)
    return {value: list(group) for value, group in groupby(ordered, key=get)}


def interleave[T](*its: Iterable[T]) -> Iterator[T]:
    """交错合并多个可迭代对象;长度不同时跳过缺失位置。

    例:list(interleave([1, 2, 3], "ab")) -> [1, "a", 2, "b", 3]
    """
    for group in zip_longest(*its, fillvalue=_MISSING):
        for item in group:
            if item is not _MISSING:
                yield item


def first_true[T](
    it: Iterable[T], pred: Callable[[T], bool], default: T | None = None
) -> T | None:
    """返回第一个让 pred 为真的元素,没有就返回 default。

    例:first_true([1, 3, 4], lambda x: x % 2 == 0) -> 4
    """
    return next((item for item in it if pred(item)), default)


def cumulative_max(nums: Iterable[float]) -> list[float]:
    """逐位记录"到目前为止的最大值"。

    例:cumulative_max([3, 1, 4, 1, 5]) -> [3, 3, 4, 4, 5]
    """
    return list(accumulate(nums, max))


def top_k_pairs(words: Iterable[str], k: int) -> list[tuple[str, int]]:
    """统计词频,返回出现次数最多的 k 个 (词, 次数),按次数降序。

    k <= 0 时返回空列表。
    例:top_k_pairs(["a", "b", "a"], 1) -> [("a", 2)]
    """
    if k <= 0:
        return []
    return Counter(words).most_common(k)


def memory_report() -> dict[str, int]:
    """对比 100 万个元素的列表推导 与 同样条件的生成器表达式的对象大小。

    返回 {"list_bytes": ..., "gen_bytes": ...},用 sys.getsizeof 测。
    例:{"list_bytes": 8448728, "gen_bytes": 224}(数字随环境略有差别)
    """
    total = 10**6
    big_list = [x for x in range(total)]
    big_gen = (x for x in range(total))
    return {
        "list_bytes": sys.getsizeof(big_list),
        "gen_bytes": sys.getsizeof(big_gen),
    }
点击展开:w1_03_decorators.py
exercises/solutions/week1/w1_03_decorators.py
"""Week 1 · M2a · 闭包与装饰器基础 参考答案

对照用:先自己把 exercises/week1/w1_03_decorators.py 写完再看这里。
验证:复制到临时目录并与 exercises/week1/test_w1_03_decorators.py
      一起运行 pytest(答案目录本身不放 test 文件,否则同名测试会冲突)。
"""

import time
from collections.abc import Callable
from functools import reduce, wraps
from typing import Any


def make_counter() -> Callable[[], int]:
    """返回一个计数器函数:每调用一次就返回比上次大 1 的整数(从 1 开始)。

    两个计数器互不影响(各自有独立的闭包变量)。
    例:c = make_counter(); c() -> 1; c() -> 2
    """
    count = 0

    def counter() -> int:
        nonlocal count
        count += 1
        return count

    return counter


def timer[**P, R](fn: Callable[P, R]) -> Callable[P, R]:
    """装饰器:打印被装饰函数的耗时(形如 "add 耗时 0.000012s"),返回原结果。

    必须用 @wraps 保住 __name__/__doc__。
    例:@timer def add(a, b): ... -> add(1, 2) 仍返回 3,并打印一行耗时
    """

    @wraps(fn)
    def wrapper(*args: P.args, **kwargs: P.kwargs) -> R:
        start = time.perf_counter()
        try:
            return fn(*args, **kwargs)
        finally:
            elapsed = time.perf_counter() - start
            print(f"{fn.__name__} 耗时 {elapsed:.6f}s")

    return wrapper


def log_calls[**P, R](
    logger_list: list[str],
) -> Callable[[Callable[P, R]], Callable[P, R]]:
    """带参装饰器:把每次调用记录成一行字符串追加到 logger_list。

    格式:位置参数与关键字参数都用 repr,形如 "add(2, 3) -> 5"、
    "power(2, exp=3) -> 8"。
    例:logs = []; @log_calls(logs) def add(a, b): return a + b
        add(2, 3) -> 5 且 logs == ["add(2, 3) -> 5"]
    """

    def decorator(fn: Callable[P, R]) -> Callable[P, R]:
        @wraps(fn)
        def wrapper(*args: P.args, **kwargs: P.kwargs) -> R:
            result = fn(*args, **kwargs)
            parts = [repr(a) for a in args]
            parts += [f"{name}={value!r}" for name, value in kwargs.items()]
            logger_list.append(f"{fn.__name__}({', '.join(parts)}) -> {result!r}")
            return result

        return wrapper

    return decorator


def retry[**P, R](
    times: int = 3,
    exceptions: tuple[type[BaseException], ...] = (ValueError,),
    delay: float = 0.0,
) -> Callable[[Callable[P, R]], Callable[P, R]]:
    """带参装饰器:调用失败就重试,最多调用 times 次,最后一次仍失败则原样抛出。

    只重试 exceptions 里列出的异常;每次重试前 sleep(delay) 秒。
    times < 1 时抛 ValueError("times 必须 >= 1")。
    例:一个"前两次抛 ValueError、第三次返回 42"的函数,被 @retry() 装饰后返回 42
    """
    if times < 1:
        raise ValueError("times 必须 >= 1")

    def decorator(fn: Callable[P, R]) -> Callable[P, R]:
        @wraps(fn)
        def wrapper(*args: P.args, **kwargs: P.kwargs) -> R:
            for attempt in range(1, times + 1):
                try:
                    return fn(*args, **kwargs)
                except exceptions:
                    if attempt == times:
                        raise
                    if delay:
                        time.sleep(delay)
            raise AssertionError("不会走到这里")

        return wrapper

    return decorator


def validate_positive[**P, R](fn: Callable[P, R]) -> Callable[P, R]:
    """装饰器:所有位置参数必须 > 0,否则抛 ValueError(关键字参数不检查)。

    例:@validate_positive def area(w, h): return w * h
        area(2, 3) -> 6;area(2, -1) -> ValueError
    """

    @wraps(fn)
    def wrapper(*args: P.args, **kwargs: P.kwargs) -> R:
        for index, value in enumerate(args, start=1):
            if isinstance(value, int | float) and value <= 0:
                raise ValueError(f"第 {index} 个位置参数必须 > 0,收到 {value!r}")
        return fn(*args, **kwargs)

    return wrapper


def once[**P, R](fn: Callable[P, R]) -> Callable[P, R]:
    """装饰器:只真正执行一次,之后不管传什么参数都返回第一次的结果。

    例:被装饰的 setup() 调用 3 次,内部只跑了 1 次
    """
    done = False
    cached: Any = None

    @wraps(fn)
    def wrapper(*args: P.args, **kwargs: P.kwargs) -> R:
        nonlocal done, cached
        if not done:
            cached = fn(*args, **kwargs)
            done = True
        return cached

    return wrapper


def compose(*fns: Callable[[Any], Any]) -> Callable[[Any], Any]:
    """从右到左组合单参函数:compose(f, g)(x) == f(g(x))。

    不传函数时返回恒等函数。
    例:compose(str, lambda x: x + 1)(1) -> "2"
    """

    def call(value: Any, fn: Callable[[Any], Any]) -> Any:
        return fn(value)

    def composed(value: Any) -> Any:
        return reduce(call, reversed(fns), value)

    return composed
点击展开:w1_04_registry_inspect.py
exercises/solutions/week1/w1_04_registry_inspect.py
"""Week 1 · M2b · 注册表与 inspect 参考答案

对照用:先自己把 exercises/week1/w1_04_registry_inspect.py 写完再看这里。
验证:复制到临时目录并与 exercises/week1/test_w1_04_registry_inspect.py
      一起运行 pytest(答案目录本身不放 test 文件,否则同名测试会冲突)。
"""

import inspect
from collections.abc import Callable
from functools import cache, singledispatch, wraps
from typing import Any, get_type_hints

# 记录两个斐波那契版本各自"真正进入函数体"的次数,测试用它对比缓存效果
CALL_COUNTS: dict[str, int] = {"naive_fib": 0, "cached_fib": 0}


class Registry:
    """名字 -> 函数 的注册表:插件、命令、工具都靠它按名字找实现。

    例:reg = Registry()
        @reg.register("add")      # 也支持不带括号:@reg.register
        def add(a, b): return a + b
        reg.get("add")(1, 2) -> 3;reg.names() -> ["add"]
    """

    def __init__(self) -> None:
        self._items: dict[str, Callable[..., Any]] = {}

    def register(
        self, name: str | Callable[..., Any] | None = None
    ) -> Callable[..., Any]:
        """注册装饰器:三种写法都要能用。

        `@reg.register`、`@reg.register()`、`@reg.register("名字")`。
        不给名字时用函数的 __name__;重名 -> ValueError("名字 add 已被注册")。
        返回值必须是原函数(这样被装饰的函数还能直接调用)。
        """
        if callable(name):  # @reg.register 不带括号:name 就是被装饰的函数
            return self._add(name.__name__, name)

        def decorator(fn: Callable[..., Any]) -> Callable[..., Any]:
            return self._add(name or fn.__name__, fn)

        return decorator

    def _add(self, name: str, fn: Callable[..., Any]) -> Callable[..., Any]:
        """真正写进字典;重名直接报错,别悄悄覆盖。"""
        if name in self._items:
            raise ValueError(f"名字 {name} 已被注册")
        self._items[name] = fn
        return fn

    def get(self, name: str) -> Callable[..., Any]:
        """按名字取函数;没有这个名字 -> KeyError("未注册的名字: xxx")。"""
        try:
            return self._items[name]
        except KeyError as e:
            raise KeyError(f"未注册的名字: {name}") from e

    def names(self) -> list[str]:
        """按注册顺序返回全部名字。"""
        return list(self._items)


def repeat[**P, R](n: int) -> Callable[[Callable[P, R]], Callable[P, list[R]]]:
    """带参装饰器:把被装饰函数连续调用 n 次,返回结果列表。

    n < 1 时抛 ValueError("n 必须 >= 1");用 @wraps 保住元数据。
    例:@repeat(3) def hi(): return "hi" -> hi() == ["hi", "hi", "hi"]
    """
    if n < 1:
        raise ValueError("n 必须 >= 1")

    def decorator(fn: Callable[P, R]) -> Callable[P, list[R]]:
        @wraps(fn)
        def wrapper(*args: P.args, **kwargs: P.kwargs) -> list[R]:
            return [fn(*args, **kwargs) for _ in range(n)]

        return wrapper

    return decorator


def naive_fib(n: int) -> int:
    """不带缓存的递归斐波那契(对照组)。

    每次进入函数体都给 CALL_COUNTS["naive_fib"] 加 1:n 稍大就会指数级爆炸。
    """
    CALL_COUNTS["naive_fib"] += 1
    if n < 2:
        return n
    return naive_fib(n - 1) + naive_fib(n - 2)


@cache
def cached_fib(n: int) -> int:
    """带缓存的递归斐波那契:同一个 n 只算一次。

    例:cached_fib(10) -> 55;cached_fib(80) 瞬间返回
    """
    CALL_COUNTS["cached_fib"] += 1
    if n < 2:
        return n
    return cached_fib(n - 1) + cached_fib(n - 2)


def _type_name(annotation: object) -> str:
    """把注解转成好读的字符串:int -> "int",没有注解 -> "Any"。"""
    if annotation is inspect.Signature.empty:
        return "Any"
    return getattr(annotation, "__name__", None) or str(annotation)


def describe(fn: Callable[..., Any]) -> dict[str, Any]:
    """反查函数签名,生成"参数说明表"(从函数生成 JSON Schema 的 Python 部分)。

    返回 {"name", "doc", "params": [{"name", "type", "default", "required"}],
    "returns"};没有注解的地方用 "Any",必填参数的 default 记 None,
    没有文档字符串时 doc 为空字符串。
    """
    signature = inspect.signature(fn)
    hints = get_type_hints(fn)
    params: list[dict[str, Any]] = []
    for name, param in signature.parameters.items():
        required = param.default is inspect.Parameter.empty
        params.append(
            {
                "name": name,
                "type": _type_name(hints.get(name, inspect.Signature.empty)),
                "default": None if required else param.default,
                "required": required,
            }
        )
    return {
        "name": fn.__name__,
        "doc": inspect.getdoc(fn) or "",
        "params": params,
        "returns": _type_name(hints.get("return", inspect.Signature.empty)),
    }


def call_with_kwargs(fn: Callable[..., Any], data: dict[str, Any]) -> Any:
    """按 fn 的签名从 data 里挑参数调用它:多余的键忽略,缺必填参数就报错。

    缺参数 -> TypeError("调用 f 缺少参数: b")(消息里要出现缺失的参数名)。
    例:def f(a, b=2): return a + b
        call_with_kwargs(f, {"a": 1, "zzz": 9}) -> 3
    """
    signature = inspect.signature(fn)
    kwargs: dict[str, Any] = {}
    missing: list[str] = []
    for name, param in signature.parameters.items():
        if name in data:
            kwargs[name] = data[name]
        elif param.default is inspect.Parameter.empty:
            missing.append(name)
    if missing:
        raise TypeError(f"调用 {fn.__name__} 缺少参数: {', '.join(missing)}")
    return fn(**kwargs)


@singledispatch
def to_str(value: object) -> str:
    """按参数类型分派的"人话描述":默认分支返回 repr(value)。

    例:to_str(3) -> "整数 3";to_str([1, 2]) -> "列表(2 项)";
        to_str("a") -> "'a'"
    """
    return repr(value)


@to_str.register
def _int_to_str(value: int) -> str:
    """int 分支。"""
    return f"整数 {value}"


@to_str.register
def _list_to_str(value: list) -> str:
    """list 分支。"""
    return f"列表({len(value)} 项)"
点击展开:w1_05_context_managers.py
exercises/solutions/week1/w1_05_context_managers.py
"""Week 1 · M3 · 上下文管理器 参考答案

对照用:先自己把 exercises/week1/w1_05_context_managers.py 写完再看这里。
验证:复制到临时目录并与 exercises/week1/test_w1_05_context_managers.py
      一起运行 pytest(答案目录本身不放 test 文件,否则同名测试会冲突)。
"""

import os
import time
from collections.abc import Iterable, Iterator
from contextlib import ExitStack, contextmanager
from pathlib import Path
from types import TracebackType
from typing import TextIO


class Timer:
    """类实现的上下文管理器:记录 with 块耗时到 self.elapsed(秒)。

    例:with Timer() as t: ...  然后 t.elapsed > 0
        块里抛异常时 elapsed 也要被记录,异常照常向外传。
    """

    def __init__(self) -> None:
        self.elapsed: float = 0.0
        self._start: float = 0.0

    def __enter__(self) -> Timer:
        """记下开始时间并返回 self(这样 `as t` 拿到的是计时器本身)。"""
        self._start = time.perf_counter()
        return self

    def __exit__(
        self,
        exc_type: type[BaseException] | None,
        exc: BaseException | None,
        tb: TracebackType | None,
    ) -> None:
        """算出耗时写进 self.elapsed;不要返回 True,异常应该继续往外抛。"""
        self.elapsed = time.perf_counter() - self._start


@contextmanager
def temp_env(**variables: str) -> Iterator[None]:
    """临时设置环境变量,退出时恢复原样(原本不存在的要删掉)。

    例:with temp_env(MODE="test"): os.environ["MODE"] == "test"
        退出后 "MODE" 又消失了;块里抛异常也要恢复。
    """
    old: dict[str, str | None] = {k: os.environ.get(k) for k in variables}
    os.environ.update(variables)
    try:
        yield
    finally:
        for key, value in old.items():
            if value is None:
                os.environ.pop(key, None)
            else:
                os.environ[key] = value


@contextmanager
def cd(path: Path) -> Iterator[Path]:
    """临时切换工作目录,退出时切回来(即使块里抛异常)。

    yield 出去的是切换后的目录(Path.cwd())。
    例:with cd(tmp_path) as here: Path.cwd() == here
    """
    before = Path.cwd()
    os.chdir(path)
    try:
        yield Path.cwd()
    finally:
        os.chdir(before)


@contextmanager
def open_all(paths: Iterable[Path]) -> Iterator[list[TextIO]]:
    """一次打开任意多个文件(UTF-8 只读),退出时全部关闭。

    例:with open_all([a, b]) as files: [f.read() for f in files]
        退出后每个 f.closed 都是 True;paths 为空时 yield 空列表。
    """
    with ExitStack() as stack:
        yield [stack.enter_context(p.open(encoding="utf-8")) for p in paths]


@contextmanager
def suppress_and_log(exc_type: type[BaseException], log: list[str]) -> Iterator[None]:
    """吞掉指定类型的异常,并把 repr(异常) 追加到 log;其它异常照常抛出。

    例:with suppress_and_log(ValueError, logs): raise ValueError("x")
        -> 不报错,logs == ["ValueError('x')"]
    """
    try:
        yield
    except exc_type as e:
        log.append(repr(e))


@contextmanager
def atomic_write(path: Path) -> Iterator[TextIO]:
    """原子写:先写同目录的 .tmp 文件,成功后 replace 覆盖目标。

    块里抛异常时目标文件保持原样,且不留下临时文件。
    例:with atomic_write(p) as f: f.write("新内容")
    """
    tmp = path.with_name(path.name + ".tmp")
    handle = tmp.open("w", encoding="utf-8")
    try:
        yield handle
    except BaseException:
        handle.close()
        tmp.unlink(missing_ok=True)
        raise
    else:
        handle.close()
        tmp.replace(path)
点击展开:w1_06_functional.py
exercises/solutions/week1/w1_06_functional.py
"""Week 1 · M3 · 函数式工具的定位 参考答案

对照用:先自己把 exercises/week1/w1_06_functional.py 写完再看这里。
验证:复制到临时目录并与 exercises/week1/test_w1_06_functional.py
      一起运行 pytest(答案目录本身不放 test 文件,否则同名测试会冲突)。
"""

from collections.abc import Callable, Iterable, Sequence
from functools import partial, reduce
from operator import itemgetter, methodcaller
from typing import Any


def sort_by_fields(
    records: Iterable[dict[str, Any]], fields: Sequence[str]
) -> list[dict[str, Any]]:
    """按多个字段升序排序,返回新列表(原列表不动)。

    fields 为空时按原顺序返回。
    例:sort_by_fields(rows, ["city", "age"]) 先按 city 再按 age
    """
    if not fields:
        return list(records)
    return sorted(records, key=itemgetter(*fields))


def compose_reduce(*fns: Callable[[Any], Any]) -> Callable[[Any], Any]:
    """用 functools.reduce 实现从右到左的函数组合。

    compose_reduce(f, g)(x) == f(g(x));不传函数时是恒等函数。
    例:compose_reduce(str, lambda x: x + 1)(1) -> "2"
    """

    def composed(value: Any) -> Any:
        return reduce(lambda acc, fn: fn(acc), reversed(fns), value)

    return composed


def squares_of_evens_map(nums: Iterable[int]) -> list[int]:
    """偶数的平方——用 map/filter 写。

    例:squares_of_evens_map([1, 2, 3, 4]) -> [4, 16]
    """
    return list(map(lambda x: x * x, filter(lambda x: x % 2 == 0, nums)))


def squares_of_evens_comprehension(nums: Iterable[int]) -> list[int]:
    """偶数的平方——用列表推导式写(结果必须与上一个函数完全一致)。

    例:squares_of_evens_comprehension([1, 2, 3, 4]) -> [4, 16]
    """
    return [x * x for x in nums if x % 2 == 0]


def _apply_rate(rate: float, amount: float) -> float:
    """按比例算出金额:保留 2 位小数。"""
    return round(amount * rate, 2)


def make_rate_applier(rate: float) -> Callable[[float], float]:
    """用 functools.partial 预填 _apply_rate 的 rate,返回只收 amount 的函数。

    例:make_rate_applier(0.1)(100) -> 10.0
    """
    return partial(_apply_rate, rate)


def upper_all(words: Iterable[str]) -> list[str]:
    """把每个字符串转大写——用 operator.methodcaller 而不是 lambda。

    例:upper_all(["ab", "cd"]) -> ["AB", "CD"]
    """
    return list(map(methodcaller("upper"), words))
点击展开:w1_07_data_model.py
exercises/solutions/week1/w1_07_data_model.py
"""Week 1 · M4a · 数据模型与魔术方法 参考答案

对照用:先自己把 exercises/week1/w1_07_data_model.py 写完再看这里。
验证:复制到临时目录并与 exercises/week1/test_w1_07_data_model.py
      一起运行 pytest(答案目录本身不放 test 文件,否则同名测试会冲突)。
"""

from collections.abc import Iterable, Iterator
from decimal import Decimal
from functools import total_ordering
from typing import Any


@total_ordering
class Money:
    """一笔钱:金额 + 货币。不同货币不能相加也不能比较。

    例:Money(Decimal("1.50"), "CNY") + Money(Decimal("0.50"), "CNY")
        -> Money(Decimal('2.00'), 'CNY')
        sum([m1, m2]) 能用(靠 __radd__ 处理起始值 0)
    """

    def __init__(self, amount: Decimal, currency: str) -> None:
        self.amount = amount
        self.currency = currency

    def __repr__(self) -> str:
        """给开发者看,最好能 eval 回来。

        例:repr(Money(Decimal("1.50"), "CNY"))
            -> "Money(Decimal('1.50'), 'CNY')"
        """
        return f"Money(Decimal('{self.amount}'), {self.currency!r})"

    def __str__(self) -> str:
        """给用户看:金额 + 空格 + 货币代码,例 "1.50 CNY"。"""
        return f"{self.amount} {self.currency}"

    def __eq__(self, other: object) -> bool:
        """金额与货币都相同才相等;和非 Money 比较返回 NotImplemented。"""
        if not isinstance(other, Money):
            return NotImplemented
        return (self.amount, self.currency) == (other.amount, other.currency)

    def __hash__(self) -> int:
        """定义了 __eq__ 就必须给 __hash__,否则对象不能进 set/dict。"""
        return hash((self.amount, self.currency))

    def __lt__(self, other: Money) -> bool:
        """同货币按金额比大小;货币不同 -> ValueError("货币不同无法比较")。"""
        if not isinstance(other, Money):
            return NotImplemented
        if self.currency != other.currency:
            raise ValueError("货币不同无法比较")
        return self.amount < other.amount

    def __add__(self, other: Money) -> Money:
        """同货币相加返回新 Money;货币不同 -> ValueError("货币不同无法相加")。"""
        if not isinstance(other, Money):
            return NotImplemented
        if self.currency != other.currency:
            raise ValueError("货币不同无法相加")
        return Money(self.amount + other.amount, self.currency)

    def __radd__(self, other: object) -> Money:
        """让 sum() 能用:sum 的起始值是 0,0 + Money 会走到这里。"""
        if other == 0:
            return self
        return NotImplemented

    def __bool__(self) -> bool:
        """金额为 0 时为假值。"""
        return self.amount != 0


class Playlist:
    """歌单:实现序列协议,`len/下标/切片/in/for/reversed` 全都能用。

    例:p = Playlist(["a", "b", "c"]);len(p) -> 3;p[0] -> "a";
        p[:2] -> Playlist(['a', 'b'])(切片返回新的 Playlist)
    """

    def __init__(self, tracks: Iterable[str]) -> None:
        self.tracks = list(tracks)

    def __repr__(self) -> str:
        """例:Playlist(['a', 'b'])。"""
        return f"Playlist({self.tracks!r})"

    def __len__(self) -> int:
        """歌曲数量。"""
        return len(self.tracks)

    def __getitem__(self, index: int | slice) -> str | Playlist:
        """整数下标返回歌名;切片返回新的 Playlist。"""
        if isinstance(index, slice):
            return Playlist(self.tracks[index])
        return self.tracks[index]

    def __contains__(self, item: object) -> bool:
        """`"a" in playlist`。"""
        return item in self.tracks

    def __iter__(self) -> Iterator[str]:
        """让 for 直接遍历歌名。"""
        return iter(self.tracks)


class Config:
    """把字典的键当属性读:cfg.host 等价于 data["host"]。

    例:Config({"host": "localhost"}).host -> "localhost"
        缺少的键 -> AttributeError("没有配置项: port")
    """

    def __init__(self, data: dict[str, Any]) -> None:
        self.data = data

    def __getattr__(self, name: str) -> Any:
        """只有正常属性查找失败时才会调用这里(self.data 不会进来)。

        找不到 -> AttributeError(f"没有配置项: {name}")。
        """
        try:
            return self.data[name]
        except KeyError as e:
            raise AttributeError(f"没有配置项: {name}") from e


class Celsius:
    """摄氏温度,支持 f-string 格式化。

    例:t = Celsius(20.0);f"{t}" -> "20.0°C";f"{t:.1f}" -> "20.0°C";
        f"{t:F}" -> "68.0°F"
    """

    def __init__(self, degrees: float) -> None:
        self.degrees = degrees

    def to_fahrenheit(self) -> float:
        """摄氏转华氏:degrees * 9 / 5 + 32。"""
        return self.degrees * 9 / 5 + 32

    def __format__(self, spec: str) -> str:
        """spec 为 "F" 时输出华氏(保留 1 位小数),否则按 spec 输出摄氏。

        例:format(Celsius(20), "F") -> "68.0°F";
            format(Celsius(20), ".2f") -> "20.00°C"
        """
        if spec == "F":
            return f"{self.to_fahrenheit():.1f}°F"
        return f"{format(self.degrees, spec)}°C"
点击展开:w1_08_dataclass_enum.py
exercises/solutions/week1/w1_08_dataclass_enum.py
"""Week 1 · M4b · dataclass 与 Enum 进阶 参考答案

对照用:先自己把 exercises/week1/w1_08_dataclass_enum.py 写完再看这里。
验证:复制到临时目录并与 exercises/week1/test_w1_08_dataclass_enum.py
      一起运行 pytest(答案目录本身不放 test 文件,否则同名测试会冲突)。
"""

from collections.abc import Iterable
from dataclasses import asdict, dataclass, field, replace
from enum import Flag, IntEnum, StrEnum, auto
from typing import Any


class Priority(IntEnum):
    """任务优先级:IntEnum 可以直接比大小、直接当整数用。

    成员:LOW = 1、NORMAL = 2、HIGH = 3。
    例:Priority.HIGH > Priority.LOW -> True;int(Priority.NORMAL) -> 2
    """

    LOW = 1
    NORMAL = 2
    HIGH = 3


class Status(StrEnum):
    """任务状态:StrEnum 成员既是枚举也是字符串,直接能进 JSON。

    成员:TODO = "todo"、DOING = "doing"、DONE = "done"。
    例:Status.DONE == "done" -> True;f"{Status.DONE}" -> "done"
    """

    TODO = "todo"
    DOING = "doing"
    DONE = "done"


class Weekday(Flag):
    """星期:Flag 可以用 | 组合、用 in 判断包含。

    例:Weekday.SAT in Weekday.WEEKEND -> True;
        (Weekday.MON | Weekday.TUE) 里有两个成员
    """

    MON = auto()
    TUE = auto()
    WED = auto()
    THU = auto()
    FRI = auto()
    SAT = auto()
    SUN = auto()
    WEEKEND = SAT | SUN


@dataclass(frozen=True, slots=True)
class Task:
    """一条任务:不可变(frozen)+ 省内存防拼错(slots)。

    例:Task("写周报") -> Task(title='写周报', priority=<Priority.NORMAL: 2>,
        status=<Status.TODO: 'todo'>, tags=[])
        Task("") -> ValueError("标题不能为空")
    """

    title: str
    priority: Priority = Priority.NORMAL
    status: Status = Status.TODO
    tags: list[str] = field(default_factory=list)

    def __post_init__(self) -> None:
        """校验:title 去掉首尾空白后不能为空,否则 ValueError("标题不能为空")。"""
        if not self.title.strip():
            raise ValueError("标题不能为空")

    def with_status(self, new: Status) -> Task:
        """返回一个只有 status 不同的新 Task(原对象不变)。

        例:Task("a").with_status(Status.DONE).status -> Status.DONE
        """
        return replace(self, status=new)


def sort_tasks(tasks: Iterable[Task]) -> list[Task]:
    """排序:优先级从高到低,同优先级按标题升序。

    例:[NORMAL "b", HIGH "a", NORMAL "a"] -> [HIGH "a", NORMAL "a", NORMAL "b"]
    """
    return sorted(tasks, key=lambda t: (-t.priority, t.title))


def to_dict(task: Task) -> dict[str, Any]:
    """把 Task 变成纯 Python 值的字典(枚举转成它的值,方便 json.dumps)。

    例:to_dict(Task("a", tags=["x"])) ->
        {"title": "a", "priority": 2, "status": "todo", "tags": ["x"]}
    """
    data = asdict(task)
    data["priority"] = task.priority.value
    data["status"] = task.status.value
    return data


def from_dict(d: dict[str, Any]) -> Task:
    """从字典还原 Task:缺省的字段用默认值,枚举从值还原。

    例:from_dict({"title": "a", "priority": 3}) ->
        Task(title='a', priority=<Priority.HIGH: 3>, ...)
        缺 title -> KeyError
    """
    return Task(
        title=d["title"],
        priority=Priority(d.get("priority", Priority.NORMAL)),
        status=Status(d.get("status", Status.TODO)),
        tags=list(d.get("tags", [])),
    )


def describe(status: Status) -> str:
    """用 match 匹配枚举成员,返回中文说明。

    Status.TODO -> "待办"、DOING -> "进行中"、DONE -> "已完成"。
    例:describe(Status.DOING) -> "进行中"
    """
    match status:
        case Status.TODO:
            return "待办"
        case Status.DOING:
            return "进行中"
        case Status.DONE:
            return "已完成"
        case _:
            raise ValueError(f"未知状态: {status!r}")
点击展开:w1_09_oop_protocols.py
exercises/solutions/week1/w1_09_oop_protocols.py
"""Week 1 · M4c · OOP 进阶 参考答案

对照用:先自己把 exercises/week1/w1_09_oop_protocols.py 写完再看这里。
验证:复制到临时目录并与 exercises/week1/test_w1_09_oop_protocols.py
      一起运行 pytest(答案目录本身不放 test 文件,否则同名测试会冲突)。
"""

from abc import ABC, abstractmethod
from collections.abc import Iterable
from math import pi
from typing import Any, Protocol, runtime_checkable


class Shape(ABC):
    """图形抽象基类:子类必须实现 area 与 perimeter。

    例:Shape() -> TypeError(抽象类不能实例化)
    """

    @abstractmethod
    def area(self) -> float:
        """面积。"""

    @abstractmethod
    def perimeter(self) -> float:
        """周长。"""

    def summary(self) -> str:
        """所有子类共用的具体方法:面积与周长各保留 2 位小数。"""
        return f"面积 {self.area():.2f} 周长 {self.perimeter():.2f}"


class Circle(Shape):
    """圆:area = pi * r²,perimeter = 2 * pi * r。

    例:Circle(1).area() -> 3.14159...;半径 <= 0 -> ValueError
    """

    def __init__(self, radius: float) -> None:
        if radius <= 0:
            raise ValueError("半径必须 > 0")
        self.radius = radius

    def area(self) -> float:
        """圆面积。"""
        return pi * self.radius**2

    def perimeter(self) -> float:
        """圆周长。"""
        return 2 * pi * self.radius


class Rect(Shape):
    """矩形:area = w * h,perimeter = 2 * (w + h)。

    例:Rect(2, 3).area() -> 6;边长 <= 0 -> ValueError
    """

    def __init__(self, width: float, height: float) -> None:
        if width <= 0 or height <= 0:
            raise ValueError("边长必须 > 0")
        self.width = width
        self.height = height

    def area(self) -> float:
        """矩形面积。"""
        return self.width * self.height

    def perimeter(self) -> float:
        """矩形周长。"""
        return 2 * (self.width + self.height)


@runtime_checkable
class Drawable(Protocol):
    """能画出自己的东西:只要有 draw() -> str 就算实现,不用继承。"""

    def draw(self) -> str:
        """返回一行字符画。"""
        ...


class Button:
    """按钮:实现 Drawable,但不继承它。

    例:Button("确定").draw() -> "[确定]"
    """

    def __init__(self, label: str) -> None:
        self.label = label

    def draw(self) -> str:
        """画成 [标签] 的样子。"""
        return f"[{self.label}]"


class Label:
    """文本标签:同样实现 Drawable,和 Button 毫无继承关系。

    例:Label("你好").draw() -> "你好"
    """

    def __init__(self, text: str) -> None:
        self.text = text

    def draw(self) -> str:
        """直接返回文本。"""
        return self.text


def render_all(items: Iterable[Drawable]) -> list[str]:
    """把每个可画对象画出来,收集成列表。

    例:render_all([Button("A"), Label("b")]) -> ["[A]", "b"]
    """
    return [item.draw() for item in items]


@runtime_checkable
class Repository(Protocol):
    """仓库接口:下周的 SQLite 版会再实现一次同样的三个方法。"""

    def add(self, key: str, value: str) -> None:
        """存一条。"""
        ...

    def get(self, key: str) -> str:
        """按 key 取,取不到抛 KeyError。"""
        ...

    def list(self) -> list[str]:
        """按插入顺序返回全部 value。"""
        ...


class InMemoryRepository:
    """用字典实现 Repository:不继承 Protocol,靠"长得像"来满足接口。

    例:r = InMemoryRepository(); r.add("a", "苹果"); r.get("a") -> "苹果"
        r.get("x") -> KeyError("找不到 x");重复 add 同一个 key 覆盖旧值
    """

    def __init__(self) -> None:
        self._items: dict[str, str] = {}

    def add(self, key: str, value: str) -> None:
        """存一条(key 重复就覆盖)。"""
        self._items[key] = value

    def get(self, key: str) -> str:
        """按 key 取;不存在 -> KeyError("找不到 x")。"""
        try:
            return self._items[key]
        except KeyError as e:
            raise KeyError(f"找不到 {key}") from e

    def list(self) -> list[str]:
        """按插入顺序返回全部 value。"""
        return list(self._items.values())


class Temperature:
    """用 @property 给属性加校验:低于绝对零度直接拒绝。

    例:t = Temperature(20); t.celsius = 25; t.fahrenheit -> 77.0
        t.celsius = -300 -> ValueError("低于绝对零度")
    """

    def __init__(self, celsius: float) -> None:
        # 走 setter,这样 __init__ 也享受校验
        self.celsius = celsius

    @property
    def celsius(self) -> float:
        """摄氏温度(读)。"""
        return self._celsius

    @celsius.setter
    def celsius(self, value: float) -> None:
        """摄氏温度(写):小于 -273.15 时 ValueError("低于绝对零度")。"""
        if value < -273.15:
            raise ValueError("低于绝对零度")
        self._celsius = value

    @property
    def fahrenheit(self) -> float:
        """只读的华氏温度:celsius * 9 / 5 + 32。"""
        return self._celsius * 9 / 5 + 32


class Plugin:
    """插件基类:每定义一个子类就自动登记到 Plugin.registry。

    例:class Hello(Plugin): ...  之后 Plugin.registry["hello"] is Hello
    """

    registry: dict[str, type[Plugin]] = {}

    def __init_subclass__(cls, **kwargs: Any) -> None:
        """定义子类时自动调用(不是实例化时)。"""
        super().__init_subclass__(**kwargs)
        Plugin.registry[cls.__name__.lower()] = cls


class Vehicle:
    """协作式 __init__ 的基类:只认自己的参数,剩下的往上传。"""

    def __init__(self, *, name: str, **kwargs: Any) -> None:
        super().__init__(**kwargs)
        self.name = name


class Car(Vehicle):
    """轿车:多一个 doors 参数(默认 4),其余参数交给父类。

    例:Car(name="A", doors=2).name -> "A"
    """

    def __init__(self, *, doors: int = 4, **kwargs: Any) -> None:
        """先 super().__init__(**kwargs) 再设自己的属性。"""
        super().__init__(**kwargs)
        self.doors = doors


class Truck(Vehicle):
    """卡车:多一个 payload_kg 参数(必填),其余参数交给父类。

    例:Truck(name="B", payload_kg=1000).payload_kg -> 1000
    """

    def __init__(self, *, payload_kg: float, **kwargs: Any) -> None:
        """同样先调 super().__init__(**kwargs),漏传 name 会直接报错。"""
        super().__init__(**kwargs)
        self.payload_kg = payload_kg
点击展开:w1_10_exceptions_syntax.py
exercises/solutions/week1/w1_10_exceptions_syntax.py
"""Week 1 · M5 · 异常进阶与 3.12–3.14 新语法(参考答案)

验证:复制到临时目录并与 exercises/week1/test_w1_10_exceptions_syntax.py 一起运行
pytest。
"""

from collections.abc import Callable, Sequence
from dataclasses import dataclass
from string.templatelib import Interpolation, Template

# ---- 任务 1:异常层次 ----


class KBError(Exception):
    """知识库所有错误的基类。调用方只需 `except KBError`。"""


class NoteNotFound(KBError):
    """按 id 找不到笔记。"""

    def __init__(self, note_id: int) -> None:
        super().__init__(f"笔记 {note_id} 不存在")
        self.note_id = note_id


class InvalidNote(KBError):
    """笔记内容不合法。"""


# ---- 任务 2:转换异常并保留原因 ----


def load_note(store: dict[int, str], note_id: int) -> str:
    """从字典取笔记;KeyError 转成 NoteNotFound,并用 `from e` 保留原因。

    例:load_note({1: "a"}, 1) -> "a";load_note({}, 9) 抛 NoteNotFound,其 __cause__
    是 KeyError
    """
    try:
        return store[note_id]
    except KeyError as e:
        raise NoteNotFound(note_id) from e


# ---- 任务 3:收集全部错误再一起抛 ----


def validate_note(text: str) -> str:
    """单条校验:去首尾空白后非空、≤ 50 字,否则抛 InvalidNote。"""
    cleaned = text.strip()
    if not cleaned:
        raise InvalidNote("笔记不能为空")
    if len(cleaned) > 50:
        raise InvalidNote(f"笔记过长:{len(cleaned)} 字")
    return cleaned


def validate_many(items: Sequence[str]) -> list[str]:
    """逐条校验,合法的收集起来;只要有错误就在最后抛 ExceptionGroup(包含全部错误)。

    例:validate_many(["a", "b"]) -> ["a", "b"]
        validate_many(["", "x" * 60]) 抛 ExceptionGroup,含 2 个 InvalidNote
    """
    ok: list[str] = []
    errors: list[Exception] = []
    for index, item in enumerate(items):
        try:
            ok.append(validate_note(item))
        except InvalidNote as e:
            e.add_note(f"第 {index} 项")
            errors.append(e)
    if errors:
        raise ExceptionGroup("部分笔记不合法", errors)
    return ok


# ---- 任务 4:match 类模式 ----


@dataclass(frozen=True)
class Add:
    x: int
    y: int


@dataclass(frozen=True)
class Greet:
    name: str


@dataclass(frozen=True)
class Quit:
    pass


type Command = Add | Greet | Quit


def handle(cmd: object) -> str:
    """用 match 的类模式分派命令。

    例:handle(Add(2, 3)) -> "结果:5";handle(Add(0, 0)) -> "零";
        handle(Greet("小明")) -> "你好,小明";handle(Quit()) -> "再见";
        其他 -> "未知命令"
    """
    match cmd:
        case Add(x=0, y=0):
            return "零"
        case Add(x=x, y=y):
            return f"结果:{x + y}"
        case Greet(name=name) if name:
            return f"你好,{name}"
        case Greet():
            return "你好,陌生人"
        case Quit():
            return "再见"
        case _:
            return "未知命令"


# ---- 任务 5:PEP 695 泛型 ----


class Stack[T]:
    """泛型栈:push/pop/peek;空栈 pop/peek 抛 IndexError。"""

    def __init__(self) -> None:
        self._items: list[T] = []

    def push(self, item: T) -> None:
        self._items.append(item)

    def pop(self) -> T:
        if not self._items:
            raise IndexError("栈为空")
        return self._items.pop()

    def peek(self) -> T:
        if not self._items:
            raise IndexError("栈为空")
        return self._items[-1]

    def __len__(self) -> int:
        return len(self._items)

    def __bool__(self) -> bool:
        return bool(self._items)


def first[T](xs: Sequence[T], default: T | None = None) -> T | None:
    """返回序列第一个元素;空序列返回 default。"""
    return xs[0] if xs else default


def apply_all[T, R](fns: Sequence[Callable[[T], R]], value: T) -> list[R]:
    """把同一个值依次交给每个函数,返回结果列表。"""
    return [fn(value) for fn in fns]


# ---- 任务 6:t-string(3.14)----


def render(t: Template) -> str:
    """渲染模板字符串:字面部分原样输出,插值部分先 format 再转义 < 和 >。

    例:name = "<b>"; render(t"hi {name}") -> "hi &lt;b&gt;"
    """
    parts: list[str] = []
    for item in t:
        if isinstance(item, Interpolation):
            text = format(item.value, item.format_spec)
            parts.append(text.replace("<", "&lt;").replace(">", "&gt;"))
        else:
            parts.append(item)
    return "".join(parts)


# ---- 任务 7:位置仅参数与 finally 语义 ----


def clamp(
    value: float, /, lo: float = 0.0, hi: float = 1.0, *, strict: bool = False
) -> float:
    """value 只能按位置传;strict=True 时越界抛 ValueError 而不是夹住。"""
    if strict and not (lo <= value <= hi):
        raise ValueError(f"{value} 不在 [{lo}, {hi}]")
    return max(lo, min(hi, value))


def read_with_cleanup(reader: Callable[[], str], log: list[str]) -> str | None:
    """演示 try/except/else/finally:reader 抛错则记录并返回 None;无论如何最后追加
    'closed'。"""
    try:
        text = reader()
    except OSError as e:
        log.append(f"error: {e}")
        return None
    else:
        log.append("ok")
        return text
    finally:
        log.append("closed")